diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..d394b52b8752f489a69888deeadf477d49143ee5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8bb9f6c25f922713aecb2a43bf2d42d559a3dbbc3ef44396e8dcf0c171a897f +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..61bd2ec0c51158a434d713ab508ff366a8ce61d5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aeca78e4ef0ca09c3c9dc1af619d8d6c62a5ad7cc2eb5e1cb43545819c3ab2e0 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a3b62bd133a7404c66cb7012f2da5ee64f93196c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d408d33f4b506621f90cc288c5605d75180856e41c97eb2027dd80bc22d772c4 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..859503d443d86e56d4e0f80bf2135951d4f81433 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7e2871b27628bb900154ef54695732d144b4d92ee183dfcd79c2e36d80feeed3 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..749e8d0ac672ea417745c99293a80966f6d61edc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:394a6f23d83467aedbce42c5ed1673db2949c93ebebcac355ea55b0341562a72 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..3388728492baf667b80729fd2c17ccaad3a18a34 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:64048e71939bae354a5cfc2dc307d22b969178dca83bd954db8b8a53bf028330 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c61878d4c3e7c59086ed63866a45be8fb8b0900d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d93683a32e9d01822aa90755c31264af9ef316e04454e8d9f0f42e5a018d880 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..39c4d89f29c97d535b4ad91b63d397a45b5d169f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9d5522748e360cabea27255f519e3f61d67f315fbd5f05216185677341a578b6 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..1ab77b67e246084cfdbfd4a817ed13ee6d1bd3f3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/0_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.759765148162842, + "learning_rate": 2e-05, + "loss": 0.1885, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.2997606992721558, + "learning_rate": 2e-05, + "loss": 0.2334, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.5568311214447021, + "learning_rate": 2e-05, + "loss": 0.0361, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.2341224104166031, + "learning_rate": 2e-05, + "loss": 0.0355, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.724339008331299, + "learning_rate": 2e-05, + "loss": 0.2466, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.4435043334960938, + "learning_rate": 2e-05, + "loss": 0.2427, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.4326348304748535, + "learning_rate": 2e-05, + "loss": 0.117, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.502455949783325, + "learning_rate": 2e-05, + "loss": 0.4245, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 5.409182548522949, + "learning_rate": 2e-05, + "loss": 0.9624, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.573906660079956, + "learning_rate": 2e-05, + "loss": 0.3429, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.15512648224830627, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.8118834495544434, + "learning_rate": 2e-05, + "loss": 0.4072, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.21215975284576416, + "learning_rate": 2e-05, + "loss": 0.1748, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.37227046489715576, + "learning_rate": 2e-05, + "loss": 0.2236, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.839745283126831, + "learning_rate": 2e-05, + "loss": 0.2748, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.12957799434661865, + "learning_rate": 2e-05, + "loss": 0.0251, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.0497076511383057, + "learning_rate": 2e-05, + "loss": 0.1551, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.895009994506836, + "learning_rate": 2e-05, + "loss": 0.3552, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.4893113076686859, + "learning_rate": 2e-05, + "loss": 0.1992, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.2643682956695557, + "learning_rate": 2e-05, + "loss": 0.0922, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.1882655769586563, + "learning_rate": 2e-05, + "loss": 0.0144, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.109602451324463, + "learning_rate": 2e-05, + "loss": 0.293, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.8853396773338318, + "learning_rate": 2e-05, + "loss": 0.1328, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.3761306703090668, + "learning_rate": 2e-05, + "loss": 0.1832, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.1210752725601196, + "learning_rate": 2e-05, + "loss": 0.074, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.428300619125366, + "learning_rate": 2e-05, + "loss": 0.0995, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.44176599383354187, + "learning_rate": 2e-05, + "loss": 0.1278, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.9758191108703613, + "learning_rate": 2e-05, + "loss": 0.1159, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 4.313375949859619, + "learning_rate": 2e-05, + "loss": 0.5219, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09026255458593369, + "learning_rate": 2e-05, + "loss": 0.0056, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7812584042549133, + "learning_rate": 2e-05, + "loss": 0.0965, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.6800825595855713, + "learning_rate": 2e-05, + "loss": 0.1061, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.4965532124042511, + "learning_rate": 2e-05, + "loss": 0.0178, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.302773118019104, + "learning_rate": 2e-05, + "loss": 0.119, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 4.1201677322387695, + "learning_rate": 2e-05, + "loss": 0.2618, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 3.332615613937378, + "learning_rate": 2e-05, + "loss": 0.1052, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 5.7126946449279785, + "learning_rate": 2e-05, + "loss": 0.9101, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.0343668460845947, + "learning_rate": 2e-05, + "loss": 0.4634, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.6162095069885254, + "learning_rate": 2e-05, + "loss": 0.279, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.5746565461158752, + "learning_rate": 2e-05, + "loss": 0.3177, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.0711337327957153, + "learning_rate": 2e-05, + "loss": 0.0829, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.581429958343506, + "learning_rate": 2e-05, + "loss": 0.1843, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.9817764163017273, + "learning_rate": 2e-05, + "loss": 0.039, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.4033353328704834, + "learning_rate": 2e-05, + "loss": 0.1229, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.7180047035217285, + "learning_rate": 2e-05, + "loss": 0.1123, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.8907074928283691, + "learning_rate": 2e-05, + "loss": 0.1681, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.818262100219727, + "learning_rate": 2e-05, + "loss": 1.121, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.6582738161087036, + "learning_rate": 2e-05, + "loss": 0.5114, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.4680819511413574, + "learning_rate": 2e-05, + "loss": 0.3985, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.24810414016246796, + "learning_rate": 2e-05, + "loss": 0.0458, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289495989583872.0, + "train_loss": 0.2355753231048584, + "train_runtime": 110.53, + "train_samples_per_second": 3.619, + "train_steps_per_second": 0.905 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289495989583872.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5e2005d8ed78197423804068c4ea9af2d6b82d14 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b43fd483b22c0b98362e12a01cbd0574094a0efb2fc57fb7b27447ff25f2c3ac +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..1cb4a1c6c856aa8c1f4713b76eb20c993a242df0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8103244c88a9cfcb174f8d9d278adc0bc0127fc3a361d35e1168688e1a79b26f +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..23776c74b696a00927d69a577ad770a5e1dfd9c2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9b6e90c54180d6764f135d50ed6915d26fdb056aca74456293bb244dfd645455 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..44a917e9d29cdc172ddec4463635f231df0b94a2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:723859a1018b828db3e4426ccfbcfa1597d88954e6a1b43f88374e6fb6bbb1d5 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..98fc841c7df52baf0feb93860ea3a5f5368f900c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:174d8bde30c02f56fe4eabdfdcaecb2adbfcb42b0be74e7d9438fd47ab79b254 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..f72dbc4a2ce046c37fc9d146ecd7c5b1c0707cbc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:166873487ba1961535ec009c0e8b1c1d8d72a93b57f7fae5922c34fa3aed4970 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..51323a0fd6cf8aca67209b99431d9a44c171221c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8a950eefd09d65bd8f3545dbe3c427d41ec5653a7bf4113dcafcf230f65620f7 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..c6d803baa7b2523d7f522527d8c427c3cb576fe6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d1914008c9c6462d95eb96625020daec333f96a21b7442e428eafff9a874dbef +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..40bf3cc05be3a17a16459c86f14060a5016d61d1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/10_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 5.540152549743652, + "learning_rate": 2e-05, + "loss": 0.4112, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.820958375930786, + "learning_rate": 2e-05, + "loss": 0.123, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.2911338806152344, + "learning_rate": 2e-05, + "loss": 0.1151, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 5.34431266784668, + "learning_rate": 2e-05, + "loss": 0.6437, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.554713726043701, + "learning_rate": 2e-05, + "loss": 0.3202, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.19470170140266418, + "learning_rate": 2e-05, + "loss": 0.0132, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 6.95194673538208, + "learning_rate": 2e-05, + "loss": 0.3958, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.7412580251693726, + "learning_rate": 2e-05, + "loss": 0.1572, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 4.699458122253418, + "learning_rate": 2e-05, + "loss": 0.1354, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.3113508224487305, + "learning_rate": 2e-05, + "loss": 0.3508, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.73491907119751, + "learning_rate": 2e-05, + "loss": 0.2418, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.238286018371582, + "learning_rate": 2e-05, + "loss": 0.1221, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.935497283935547, + "learning_rate": 2e-05, + "loss": 0.5967, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 9.50597858428955, + "learning_rate": 2e-05, + "loss": 0.6172, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 8.673219680786133, + "learning_rate": 2e-05, + "loss": 0.5515, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.0694637298583984, + "learning_rate": 2e-05, + "loss": 0.0626, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.6291310787200928, + "learning_rate": 2e-05, + "loss": 0.1058, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.465129554271698, + "learning_rate": 2e-05, + "loss": 0.2965, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.1682217121124268, + "learning_rate": 2e-05, + "loss": 0.176, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.369879961013794, + "learning_rate": 2e-05, + "loss": 0.2682, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.5110698938369751, + "learning_rate": 2e-05, + "loss": 0.0655, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.19209489226341248, + "learning_rate": 2e-05, + "loss": 0.5971, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.3685367703437805, + "learning_rate": 2e-05, + "loss": 0.0502, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.7339094877243042, + "learning_rate": 2e-05, + "loss": 0.2632, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.006829261779785, + "learning_rate": 2e-05, + "loss": 0.1007, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.7255167961120605, + "learning_rate": 2e-05, + "loss": 0.9134, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 12.463668823242188, + "learning_rate": 2e-05, + "loss": 1.3237, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.31620460748672485, + "learning_rate": 2e-05, + "loss": 0.0115, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 4.888059616088867, + "learning_rate": 2e-05, + "loss": 0.1024, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 9.215259552001953, + "learning_rate": 2e-05, + "loss": 0.335, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 6.971097469329834, + "learning_rate": 2e-05, + "loss": 0.4281, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.5247466564178467, + "learning_rate": 2e-05, + "loss": 0.0873, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.305556297302246, + "learning_rate": 2e-05, + "loss": 0.8254, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.3106958866119385, + "learning_rate": 2e-05, + "loss": 0.4049, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.9962230324745178, + "learning_rate": 2e-05, + "loss": 0.0339, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 15.942362785339355, + "learning_rate": 2e-05, + "loss": 0.917, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.5281805992126465, + "learning_rate": 2e-05, + "loss": 0.2087, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 9.443400382995605, + "learning_rate": 2e-05, + "loss": 0.4157, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 6.793636798858643, + "learning_rate": 2e-05, + "loss": 0.8669, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.395467758178711, + "learning_rate": 2e-05, + "loss": 0.732, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 7.989163398742676, + "learning_rate": 2e-05, + "loss": 0.402, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.12291645258665085, + "learning_rate": 2e-05, + "loss": 0.0088, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 5.78603982925415, + "learning_rate": 2e-05, + "loss": 0.3615, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.31522274017334, + "learning_rate": 2e-05, + "loss": 0.5243, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.1642500162124634, + "learning_rate": 2e-05, + "loss": 0.1487, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.8082680702209473, + "learning_rate": 2e-05, + "loss": 0.4324, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.412454605102539, + "learning_rate": 2e-05, + "loss": 0.7421, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.5595290660858154, + "learning_rate": 2e-05, + "loss": 0.121, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.6381373405456543, + "learning_rate": 2e-05, + "loss": 0.2491, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.7060370445251465, + "learning_rate": 2e-05, + "loss": 0.5233, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2220576018006016.0, + "train_loss": 0.35800392150878907, + "train_runtime": 82.3263, + "train_samples_per_second": 4.859, + "train_steps_per_second": 1.215 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2220576018006016.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..9724ea2b7428cb70d3b6ff23669f007f286d32ac --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9c7dabf9a1e6938c5ad868400abd8bfb93de9b8018beb56e924aee4ff8da4615 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..88e3cb71e70f83777fe7be491bcdf0541c7f6b4d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0dbfa77c2b17796f76a16e1bc345aa3003103d7751c7f027622bb97b93371097 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..8d7456e98ac648d6f591c12a178cb2c3d96f01e7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ec59bb41278eef1a16c3ec9c38b5725a99e3a0d1a8a0744597d6d53665e49e56 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0f756c2b71feb72aadab69ba09c4ed320a60201c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba3d9536f47b9667a637006fb5cd0c09495296df2b95c762598639f0a28211e6 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..8481a99b3ba949e0a2d567cc2ca069b82311fed1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8416029955098733bd16102f51e5190913435a273688a56109d7e28971a62703 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6ce8cf24b1054fe0adc063baaeca28f907a3507f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:512ce8d54f2e758e8aaff9a3225b056dd058a92c7d661054042a8d2f964b6a48 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b0177ba6283cf925a43c78806fad34491e2ed3d2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:df17747c40ec198afdcc951a4e99fd8683cd84e079d089eadb40ca3dbb8ca8e6 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..710a621ba082efbdcb321855dd486de651cf6dd3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3eaf6003474ef273ca56a9a823bc0540d3448b4f0f8be5e9c93504a9ff99277b +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..8e2883c98ad6c80a504f465360080022c9d4a8f7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/11_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 5.755228042602539, + "learning_rate": 2e-05, + "loss": 0.4839, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.927358150482178, + "learning_rate": 2e-05, + "loss": 0.3677, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.49658465385437, + "learning_rate": 2e-05, + "loss": 0.381, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.028351306915283, + "learning_rate": 2e-05, + "loss": 0.2484, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.260554313659668, + "learning_rate": 2e-05, + "loss": 0.793, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.9733984470367432, + "learning_rate": 2e-05, + "loss": 0.4696, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 5.129627227783203, + "learning_rate": 2e-05, + "loss": 0.9727, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.833235263824463, + "learning_rate": 2e-05, + "loss": 0.71, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.3442604541778564, + "learning_rate": 2e-05, + "loss": 0.5811, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.6310151815414429, + "learning_rate": 2e-05, + "loss": 0.359, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.7572506666183472, + "learning_rate": 2e-05, + "loss": 0.2868, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.0499885082244873, + "learning_rate": 2e-05, + "loss": 0.4556, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.5038318634033203, + "learning_rate": 2e-05, + "loss": 0.3833, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 7.785128593444824, + "learning_rate": 2e-05, + "loss": 0.4614, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 4.476345062255859, + "learning_rate": 2e-05, + "loss": 0.54, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.5807504653930664, + "learning_rate": 2e-05, + "loss": 0.4373, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.6199612617492676, + "learning_rate": 2e-05, + "loss": 0.5078, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 5.068996906280518, + "learning_rate": 2e-05, + "loss": 0.4609, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.555264949798584, + "learning_rate": 2e-05, + "loss": 0.4846, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.7653462290763855, + "learning_rate": 2e-05, + "loss": 0.4608, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 6.7089691162109375, + "learning_rate": 2e-05, + "loss": 0.3584, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 5.8125996589660645, + "learning_rate": 2e-05, + "loss": 0.6569, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.4853155612945557, + "learning_rate": 2e-05, + "loss": 0.4939, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.0506442785263062, + "learning_rate": 2e-05, + "loss": 0.2601, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 5.831657886505127, + "learning_rate": 2e-05, + "loss": 0.6602, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.004698753356934, + "learning_rate": 2e-05, + "loss": 0.5503, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.582853078842163, + "learning_rate": 2e-05, + "loss": 0.3324, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 4.086846828460693, + "learning_rate": 2e-05, + "loss": 0.2622, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.316006898880005, + "learning_rate": 2e-05, + "loss": 0.235, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.871242880821228, + "learning_rate": 2e-05, + "loss": 0.3149, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.0035479068756104, + "learning_rate": 2e-05, + "loss": 0.3652, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.884969711303711, + "learning_rate": 2e-05, + "loss": 0.3341, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 6.4319562911987305, + "learning_rate": 2e-05, + "loss": 0.6167, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.11548277735710144, + "learning_rate": 2e-05, + "loss": 0.1651, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.814781427383423, + "learning_rate": 2e-05, + "loss": 0.3645, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 3.4434401988983154, + "learning_rate": 2e-05, + "loss": 0.4271, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.981860876083374, + "learning_rate": 2e-05, + "loss": 0.5566, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.0367769002914429, + "learning_rate": 2e-05, + "loss": 0.4013, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 6.169937610626221, + "learning_rate": 2e-05, + "loss": 0.6856, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 4.215674877166748, + "learning_rate": 2e-05, + "loss": 0.6436, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.0360214710235596, + "learning_rate": 2e-05, + "loss": 0.7363, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 4.0028862953186035, + "learning_rate": 2e-05, + "loss": 0.6914, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.7605392336845398, + "learning_rate": 2e-05, + "loss": 0.3038, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.8313522338867188, + "learning_rate": 2e-05, + "loss": 0.4746, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.02323991432785988, + "learning_rate": 2e-05, + "loss": 0.4973, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 5.001913070678711, + "learning_rate": 2e-05, + "loss": 0.3467, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.9565664529800415, + "learning_rate": 2e-05, + "loss": 0.4761, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.934812307357788, + "learning_rate": 2e-05, + "loss": 0.4241, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.334172010421753, + "learning_rate": 2e-05, + "loss": 0.3789, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 7.261927604675293, + "learning_rate": 2e-05, + "loss": 0.4741, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2191347477905408.0, + "train_loss": 0.4666449117660523, + "train_runtime": 64.5127, + "train_samples_per_second": 6.2, + "train_steps_per_second": 1.55 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2191347477905408.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..97322a6c98c974944632c997e8b9dbf2a82dd2fa --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6d4afee4d1a1c8be3faf3ec0f2f95a54a22205e4c00c40ee6186d5456d64fb8c +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..49f7b87db6eb0f0eb14e69b7aa4c0165053286ce --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:65f6deb5ebd27fae9647a6aa87cf2025b7df287976fe43f55d8659a0a3ec8206 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..6ecfccf300062a71f2ab9a1028727b87ab01ee5f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:acf3b7b44d41068d7ec41829ea125b4836111669f344efe3b993d9d58829e674 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..6c3b2b810cad8c035b27297696dce15419e556f0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e6c9fdf5d2bdd6f772d243ef751ff09df048c6154ab0b54061ba3b3698e51c93 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..0d019e698e02f598ab96667be7bc361d5f57fa27 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f1215abb22b5f3080310503750e09f5540b438773e81753bd7a1be135fb7ea31 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..842738bcb8e08673cc6b6dd3b6977a4015fc9a99 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8fa0b2d3b8067f842b2ac6fdde666556f9f0ad2f82fa22823a8f92213eb015a +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..10eae58e339e07605d8e5e6ee7c55f94ed07ca4a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ccfd3fe6c7e97d0404cb0e743cc1812b686fe243246905373d7ab085be433fc8 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..2148490f98005d317865277692dc9274425aecc6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4589a160b67e1c081ab2e0d17dbcdead79c6ba28c25b2e158ab6f6f3baf883b2 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..dc13bb39423285e002af8f4c18dcc088ed077eab --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/12_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.2920405864715576, + "learning_rate": 2e-05, + "loss": 0.0386, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.0024220957420766354, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.042690176516771317, + "learning_rate": 2e-05, + "loss": 0.0126, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.620497465133667, + "learning_rate": 2e-05, + "loss": 0.0201, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.22719046473503113, + "learning_rate": 2e-05, + "loss": 0.0145, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.07465826719999313, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 4.26301908493042, + "learning_rate": 2e-05, + "loss": 0.604, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.5791594982147217, + "learning_rate": 2e-05, + "loss": 0.2624, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.459093302488327, + "learning_rate": 2e-05, + "loss": 0.1353, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.221980571746826, + "learning_rate": 2e-05, + "loss": 0.2825, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.17288897931575775, + "learning_rate": 2e-05, + "loss": 0.0922, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.26775988936424255, + "learning_rate": 2e-05, + "loss": 0.0136, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.4521391987800598, + "learning_rate": 2e-05, + "loss": 0.0704, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.4336092472076416, + "learning_rate": 2e-05, + "loss": 0.1329, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.026113972067832947, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.06741499155759811, + "learning_rate": 2e-05, + "loss": 0.0105, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.10180412232875824, + "learning_rate": 2e-05, + "loss": 0.0316, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.2327575832605362, + "learning_rate": 2e-05, + "loss": 0.0671, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.21509437263011932, + "learning_rate": 2e-05, + "loss": 0.0205, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.175586462020874, + "learning_rate": 2e-05, + "loss": 0.0603, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 4.097664833068848, + "learning_rate": 2e-05, + "loss": 0.3459, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.04044518247246742, + "learning_rate": 2e-05, + "loss": 0.0197, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.225965976715088, + "learning_rate": 2e-05, + "loss": 0.1946, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.1212005615234375, + "learning_rate": 2e-05, + "loss": 0.0359, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.05873919650912285, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.07613237202167511, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.06085467338562012, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.0594962015748024, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.022673513740301132, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09007851779460907, + "learning_rate": 2e-05, + "loss": 0.0046, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.006416500546038151, + "learning_rate": 2e-05, + "loss": 0.0075, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.008369038812816143, + "learning_rate": 2e-05, + "loss": 0.0033, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.07328344881534576, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7168121933937073, + "learning_rate": 2e-05, + "loss": 0.0257, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.006075312849134207, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.027319282293319702, + "learning_rate": 2e-05, + "loss": 0.3296, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.05318591743707657, + "learning_rate": 2e-05, + "loss": 0.0021, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.008996729739010334, + "learning_rate": 2e-05, + "loss": 0.0185, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.26109060645103455, + "learning_rate": 2e-05, + "loss": 0.2289, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.0030557974241673946, + "learning_rate": 2e-05, + "loss": 0.0078, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.0699257180094719, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.06251349300146103, + "learning_rate": 2e-05, + "loss": 0.0082, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0034623832907527685, + "learning_rate": 2e-05, + "loss": 0.0086, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.582453727722168, + "learning_rate": 2e-05, + "loss": 0.2078, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.005348577164113522, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.48420047760009766, + "learning_rate": 2e-05, + "loss": 0.0835, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.969717025756836, + "learning_rate": 2e-05, + "loss": 0.4384, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.030660836026072502, + "learning_rate": 2e-05, + "loss": 0.3863, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.05689483880996704, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.1647011041641235, + "learning_rate": 2e-05, + "loss": 0.0288, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5285302658662400.0, + "train_loss": 0.08597439825534821, + "train_runtime": 99.8483, + "train_samples_per_second": 4.006, + "train_steps_per_second": 1.002 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5285302658662400.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..99e8e93ea44263422990c1ac39834a75326b2340 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b6045ff56789a2787a3639b748e163b57c539b585db7f974197cb4d2616f738a +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ff73530204580e06064aa0c6ac1af40834f233df --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2427fc2ce472d6b450bc333eea0aaa7ccba73863e4570b5e475ba2ed62a2674 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..7baf7487075e16021961792a2af5b71b7614c36d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:abbe7751f467131dbb9b0ed8e7049d134839bcda9e91f70ccb1360425ea95344 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e3fac92b04084a9771bf788a2374450a0bf3df43 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:54b1a1314d2b743e3c567022ffc38e6db93d7a89db83083d75dc0ad60020b818 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..6be081312347004f11f86e41e4ba4d1193697207 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e198c8ca1c2f53cdf9ac3cda0051878b7582318d64562b4c88545fa94ec4647b +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..81d2b45831b95fbf6d3526060fa3792a932c246b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dd3b9f4ca4475c87bc11a49a5debfc1ae0450831f2d55196a4cc9e31e1ea89f7 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..47436d1a312087e5d9718c71e567a8334ee11fb6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb14a353f3c12858de38a9f49f39ca99e3ce1f0800729d70182c9241f4c48689 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..dab470a89c4b45dab80e5196c019c401e3f2bb98 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:33d113484917d92fca14d16bf0d06ce34c11e46aa6d5179c8dfc8ce893d1f19d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..6e33857afda5daac6e7e680026c5c8c4d4be5a5d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/13_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.778566360473633, + "learning_rate": 2e-05, + "loss": 0.3427, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.5951638221740723, + "learning_rate": 2e-05, + "loss": 0.2446, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.247312307357788, + "learning_rate": 2e-05, + "loss": 0.3885, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.3595407009124756, + "learning_rate": 2e-05, + "loss": 0.0998, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8933051228523254, + "learning_rate": 2e-05, + "loss": 0.0959, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.20830616354942322, + "learning_rate": 2e-05, + "loss": 0.0334, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.4658172130584717, + "learning_rate": 2e-05, + "loss": 0.1042, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 4.962435722351074, + "learning_rate": 2e-05, + "loss": 0.429, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.4839186668395996, + "learning_rate": 2e-05, + "loss": 0.1982, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.563016414642334, + "learning_rate": 2e-05, + "loss": 1.2359, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 6.918004989624023, + "learning_rate": 2e-05, + "loss": 0.5588, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 7.100998878479004, + "learning_rate": 2e-05, + "loss": 0.2685, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.131011009216309, + "learning_rate": 2e-05, + "loss": 0.9571, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.1804964542388916, + "learning_rate": 2e-05, + "loss": 0.1269, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.373103380203247, + "learning_rate": 2e-05, + "loss": 0.2773, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.8673224449157715, + "learning_rate": 2e-05, + "loss": 0.1382, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.9214558601379395, + "learning_rate": 2e-05, + "loss": 0.2176, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.2953531742095947, + "learning_rate": 2e-05, + "loss": 0.3478, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.48850154876709, + "learning_rate": 2e-05, + "loss": 0.1105, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.431159257888794, + "learning_rate": 2e-05, + "loss": 0.2718, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.0679343119263649, + "learning_rate": 2e-05, + "loss": 0.2713, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.108593463897705, + "learning_rate": 2e-05, + "loss": 0.3638, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.01567334122955799, + "learning_rate": 2e-05, + "loss": 0.0121, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.7148890495300293, + "learning_rate": 2e-05, + "loss": 0.2139, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.34122204780578613, + "learning_rate": 2e-05, + "loss": 0.1943, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.237427234649658, + "learning_rate": 2e-05, + "loss": 0.4129, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.882807493209839, + "learning_rate": 2e-05, + "loss": 0.3954, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.5616434812545776, + "learning_rate": 2e-05, + "loss": 0.1976, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.07391391694545746, + "learning_rate": 2e-05, + "loss": 0.0307, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.12220025807619095, + "learning_rate": 2e-05, + "loss": 0.0502, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.797428846359253, + "learning_rate": 2e-05, + "loss": 0.184, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.5082601308822632, + "learning_rate": 2e-05, + "loss": 0.067, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.5079703330993652, + "learning_rate": 2e-05, + "loss": 0.3816, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 4.023115158081055, + "learning_rate": 2e-05, + "loss": 0.3118, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.5129783153533936, + "learning_rate": 2e-05, + "loss": 0.2305, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.47571924328804016, + "learning_rate": 2e-05, + "loss": 0.0345, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.9408183693885803, + "learning_rate": 2e-05, + "loss": 0.563, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.6250778436660767, + "learning_rate": 2e-05, + "loss": 0.1207, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.6771740913391113, + "learning_rate": 2e-05, + "loss": 0.2536, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.673983097076416, + "learning_rate": 2e-05, + "loss": 0.0871, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.856926441192627, + "learning_rate": 2e-05, + "loss": 0.1238, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 6.238300323486328, + "learning_rate": 2e-05, + "loss": 0.7711, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 6.600960731506348, + "learning_rate": 2e-05, + "loss": 0.6233, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.1156471967697144, + "learning_rate": 2e-05, + "loss": 0.135, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.041778326034546, + "learning_rate": 2e-05, + "loss": 0.1254, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.35611891746521, + "learning_rate": 2e-05, + "loss": 0.5919, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.3482422828674316, + "learning_rate": 2e-05, + "loss": 0.0314, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.902723789215088, + "learning_rate": 2e-05, + "loss": 0.2435, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.3907882273197174, + "learning_rate": 2e-05, + "loss": 0.3463, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.294113636016846, + "learning_rate": 2e-05, + "loss": 0.3237, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5322110830379008.0, + "train_loss": 0.2827692413330078, + "train_runtime": 110.5433, + "train_samples_per_second": 3.618, + "train_steps_per_second": 0.905 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5322110830379008.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..68f6305d4739966cbe66030280f88dd738f598dc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:11b97ed5f57a071bb8d6d641af5ccd55ab115b7ba5ba65cef3293e05f4a0740e +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..17f1533c595bfe3a386da553c55eb5152bf4d8f3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e82666a5ee5251162c7a298a5785069be482719070acffabd09a3bea7951d9e8 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..db35126b95b44f19549eeaeae57333659d99e4f5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3a61109378eb3292593745ee8b3477e53e54e11da2a1e7de078f8ab0ab06d439 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..33537ec4aaae7cc93522954c5a396579339139ce --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:de2fa8c21d1b8a2894e8b4e5c1498375bd2be76dc3837828926973ea44ab50c3 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..c72044dc39daae2c17f15039cbf206948f4d531b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d6b3863d8225a6a2c32460c2ca425e71237b4f897dc75dc369e7f9e807eae9df +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d3ee90106befbef24ee2fe727c184ad0e83c5425 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf97e648a23af90a7d62633a2b272f74db0f6f17566c883f81997199b4a42243 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2075eaac0dfdefd4054e1de6aa5560569715cf4b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e3e041742868b78d5e3a4275a72ff511385306af3307f351447a562f62b83d49 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..3695d944968280b2d5747abe07e6a0d3a1108d82 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09c43689fc704fb46da00021f500f6a558fb6b908165620ede33643a63aaec33 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..5d34293627f674bc3019e6117e0c6b93edff4d5e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/14_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.6643385887145996, + "learning_rate": 2e-05, + "loss": 0.4654, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.549673318862915, + "learning_rate": 2e-05, + "loss": 0.3132, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.30682796239852905, + "learning_rate": 2e-05, + "loss": 0.0685, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.030418802052736282, + "learning_rate": 2e-05, + "loss": 0.0014, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.666572570800781, + "learning_rate": 2e-05, + "loss": 0.6256, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.6245814561843872, + "learning_rate": 2e-05, + "loss": 0.2397, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.417027235031128, + "learning_rate": 2e-05, + "loss": 0.3518, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.02365238033235073, + "learning_rate": 2e-05, + "loss": 0.0039, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.415682077407837, + "learning_rate": 2e-05, + "loss": 0.1614, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.3901709318161011, + "learning_rate": 2e-05, + "loss": 0.03, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.15750576555728912, + "learning_rate": 2e-05, + "loss": 0.0296, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.018209584057331085, + "learning_rate": 2e-05, + "loss": 0.0015, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.6798834800720215, + "learning_rate": 2e-05, + "loss": 0.5891, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 7.348636150360107, + "learning_rate": 2e-05, + "loss": 0.3555, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.469543695449829, + "learning_rate": 2e-05, + "loss": 0.2062, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.7010553479194641, + "learning_rate": 2e-05, + "loss": 0.0688, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.16325834393501282, + "learning_rate": 2e-05, + "loss": 0.1559, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.6043805480003357, + "learning_rate": 2e-05, + "loss": 0.0346, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.1538206934928894, + "learning_rate": 2e-05, + "loss": 0.0165, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.6587228775024414, + "learning_rate": 2e-05, + "loss": 0.3646, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.1928708851337433, + "learning_rate": 2e-05, + "loss": 0.1483, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.559927225112915, + "learning_rate": 2e-05, + "loss": 0.3288, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.7576311826705933, + "learning_rate": 2e-05, + "loss": 0.0379, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.173165202140808, + "learning_rate": 2e-05, + "loss": 0.0556, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.03007105365395546, + "learning_rate": 2e-05, + "loss": 0.1686, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.027469219639897346, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.1872060298919678, + "learning_rate": 2e-05, + "loss": 0.04, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.792096734046936, + "learning_rate": 2e-05, + "loss": 0.1033, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.6836594343185425, + "learning_rate": 2e-05, + "loss": 0.1692, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.025690946727991104, + "learning_rate": 2e-05, + "loss": 0.0029, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7871533036231995, + "learning_rate": 2e-05, + "loss": 0.1753, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.6226503252983093, + "learning_rate": 2e-05, + "loss": 0.0291, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.0744810104370117, + "learning_rate": 2e-05, + "loss": 0.1295, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.7682394981384277, + "learning_rate": 2e-05, + "loss": 0.2178, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.318707674741745, + "learning_rate": 2e-05, + "loss": 0.1649, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.10394205898046494, + "learning_rate": 2e-05, + "loss": 0.0228, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.8076616525650024, + "learning_rate": 2e-05, + "loss": 0.1041, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.129563570022583, + "learning_rate": 2e-05, + "loss": 0.132, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.6246461272239685, + "learning_rate": 2e-05, + "loss": 0.1212, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.8219376802444458, + "learning_rate": 2e-05, + "loss": 0.0555, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.5030476450920105, + "learning_rate": 2e-05, + "loss": 0.0618, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.7511129975318909, + "learning_rate": 2e-05, + "loss": 0.061, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.6618772149085999, + "learning_rate": 2e-05, + "loss": 0.023, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.4337066113948822, + "learning_rate": 2e-05, + "loss": 0.0292, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.53242826461792, + "learning_rate": 2e-05, + "loss": 0.0621, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.862425684928894, + "learning_rate": 2e-05, + "loss": 0.0433, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.14855121076107025, + "learning_rate": 2e-05, + "loss": 0.0129, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.7792656421661377, + "learning_rate": 2e-05, + "loss": 0.0381, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.26111164689064026, + "learning_rate": 2e-05, + "loss": 0.0153, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.005244841333478689, + "learning_rate": 2e-05, + "loss": 0.0376, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5337401710870528.0, + "train_loss": 0.13353644847869872, + "train_runtime": 95.1987, + "train_samples_per_second": 4.202, + "train_steps_per_second": 1.05 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5337401710870528.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..cac44046a262258d126657a88ee8f78f95a0832b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e4f40ad5ed185a601fc74356a86cfd2f7de1a28736f013daebbcdc6f52548b4 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..50ea8bc8a25b490c9ed3109d76747106717253f7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8fb4f1456e6c9abb9b6bf17a99f42541e2505e23b2529061c8c802714b4ffdd +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..979992ff2fb2743a3968f03cf1415b43aa0672cd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:38536aad0ed4c38b4be705a0bad2a79ab1522ea6b7e20c6e94a38657ae3211db +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..ccd785666495c91b435cf498a4884a6bffde8274 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:30cf6fe3820ebb98d3bfbe05519643e92dcab7c307e965d291824fb85c03cb42 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..c376575b9586f2006bda0e03eba3c53d63b294ee --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:85463d825a596336b16392ff636a1e561250a21c55183de0f2466e1ff0630d5e +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..3dcd15e326b81fa3f5cb91230e2a59eaac1931cf --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:723e00af74d4cec96b922d864b46232024f42ab093d5a8fcabb4f7ee2074efdc +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..4279ca82aaea187714b5917e284f20463546d24f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:57ef88689fc6a4984cbee4f810265c5f11ceda81d22133065cb12419b5cde04d +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..32e478235a7966619e0f397f68a0d90e6c40a801 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eed38ad8443b74efda04fc1148478447aee44fe6908ab85867d897e9d37862ec +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..5088b03d3b8ab46fda26fc117bc2e705224f2676 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/15_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.993752956390381, + "learning_rate": 2e-05, + "loss": 0.134, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 7.597855567932129, + "learning_rate": 2e-05, + "loss": 0.6713, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 5.59683895111084, + "learning_rate": 2e-05, + "loss": 0.5211, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.9701704978942871, + "learning_rate": 2e-05, + "loss": 0.105, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.9076824188232422, + "learning_rate": 2e-05, + "loss": 0.0644, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.001618504524231, + "learning_rate": 2e-05, + "loss": 0.1731, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.8202154040336609, + "learning_rate": 2e-05, + "loss": 0.0347, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.4118986129760742, + "learning_rate": 2e-05, + "loss": 0.3755, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.4547221660614014, + "learning_rate": 2e-05, + "loss": 0.0672, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.3161427974700928, + "learning_rate": 2e-05, + "loss": 0.2075, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 5.867923259735107, + "learning_rate": 2e-05, + "loss": 0.4458, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 4.831803321838379, + "learning_rate": 2e-05, + "loss": 0.3044, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.863537311553955, + "learning_rate": 2e-05, + "loss": 0.3537, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.1004350185394287, + "learning_rate": 2e-05, + "loss": 0.3101, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.9127001762390137, + "learning_rate": 2e-05, + "loss": 0.233, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.4059998989105225, + "learning_rate": 2e-05, + "loss": 0.0785, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.185575008392334, + "learning_rate": 2e-05, + "loss": 0.3879, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.32736799120903015, + "learning_rate": 2e-05, + "loss": 0.0521, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.377376556396484, + "learning_rate": 2e-05, + "loss": 0.3225, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.5486228466033936, + "learning_rate": 2e-05, + "loss": 0.1677, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.2775516211986542, + "learning_rate": 2e-05, + "loss": 0.0667, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.789041757583618, + "learning_rate": 2e-05, + "loss": 0.4739, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 8.830492973327637, + "learning_rate": 2e-05, + "loss": 0.1451, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.9674887657165527, + "learning_rate": 2e-05, + "loss": 0.0588, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 5.155989646911621, + "learning_rate": 2e-05, + "loss": 0.6953, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.765751838684082, + "learning_rate": 2e-05, + "loss": 0.2855, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.12423150986433029, + "learning_rate": 2e-05, + "loss": 0.0169, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.8664994835853577, + "learning_rate": 2e-05, + "loss": 0.0308, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.0248830318450928, + "learning_rate": 2e-05, + "loss": 0.2541, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.071001648902893, + "learning_rate": 2e-05, + "loss": 0.1032, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 9.72658920288086, + "learning_rate": 2e-05, + "loss": 0.5628, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 5.664763450622559, + "learning_rate": 2e-05, + "loss": 0.1822, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.7218841314315796, + "learning_rate": 2e-05, + "loss": 0.0912, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 6.914689064025879, + "learning_rate": 2e-05, + "loss": 0.4173, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.9764505624771118, + "learning_rate": 2e-05, + "loss": 0.5121, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 3.8638393878936768, + "learning_rate": 2e-05, + "loss": 0.2675, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.6353455781936646, + "learning_rate": 2e-05, + "loss": 0.2072, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 4.367946624755859, + "learning_rate": 2e-05, + "loss": 0.2816, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.0566858053207397, + "learning_rate": 2e-05, + "loss": 0.0441, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.011694312095642, + "learning_rate": 2e-05, + "loss": 0.1609, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.220724582672119, + "learning_rate": 2e-05, + "loss": 0.1802, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.271707773208618, + "learning_rate": 2e-05, + "loss": 0.1697, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.033854007720947, + "learning_rate": 2e-05, + "loss": 0.6383, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 7.362668514251709, + "learning_rate": 2e-05, + "loss": 0.7583, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.794032573699951, + "learning_rate": 2e-05, + "loss": 0.3075, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.7548731565475464, + "learning_rate": 2e-05, + "loss": 0.2403, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.366307258605957, + "learning_rate": 2e-05, + "loss": 0.0616, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.46035662293434143, + "learning_rate": 2e-05, + "loss": 0.8321, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.098369598388672, + "learning_rate": 2e-05, + "loss": 0.238, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.764628291130066, + "learning_rate": 2e-05, + "loss": 0.0894, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2213537237696512.0, + "train_loss": 0.26763219833374025, + "train_runtime": 63.2319, + "train_samples_per_second": 6.326, + "train_steps_per_second": 1.581 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2213537237696512.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..daaba5910bef826b1cf881df212fe7a1fe7cedda --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d494a2bd7e4a23e3850ebbe291e8bae6eb17be514bab298d3a507852ff2e6fc1 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..e23c295839fd36ae26c95059a0dc8bbbbeff2fb4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3cb17fb566ef2e8e6efb50b6ba1956d17909702b5497ff7fcb51de8c005b8998 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..beaffc06f496a40cdc5fb4c22f65e573d0569971 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d9a0353038f56cc75335e5dfef0c03bba59c8a9ed0fbb4b8b5230f12476e72b0 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..45d4bd4b2f9b830e428900fc17f9964e5959d879 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f1880ec3f2adde719862815bca643aa9873c31b71310b80bd8a7deba4ea9b0b4 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..3b65e9d3d6baa63b62deec0f485b9655377dcfd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2108d0cecb0ef15b69c4f561f33df1ed3923263bd73db43092dc564e8cbd5004 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..a93f60de99e61e3e9f0db565f3c57cea21c13c4b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:17f2f4d11902c44181f3ee35aef5040f822fab1f793f2f1ee787dac6bccf4686 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..d6d1d016ffad80b53d4497afe86cec3f9e3f99b3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4e8b7bd47717af0ffc9fef834b339dcd38923fd274746a2159104ddced1f3e49 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..3d52e2cfeb9ac246c0b44ff0ddda09a8cf5b1ae0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebd7834614c511dc838622accb132c1b8cd272578d2fc65cb46f01ab740dd549 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..20a829d4e9d7f4a4966df5f7aff54707df8afbd0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/16_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.924595832824707, + "learning_rate": 2e-05, + "loss": 0.3051, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.33572396636009216, + "learning_rate": 2e-05, + "loss": 0.0364, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 10.899779319763184, + "learning_rate": 2e-05, + "loss": 0.8442, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 5.387420654296875, + "learning_rate": 2e-05, + "loss": 0.2751, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.831040859222412, + "learning_rate": 2e-05, + "loss": 0.2829, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.09246519207954407, + "learning_rate": 2e-05, + "loss": 0.0484, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.6267755031585693, + "learning_rate": 2e-05, + "loss": 0.1533, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.214953899383545, + "learning_rate": 2e-05, + "loss": 0.3068, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.0170130729675293, + "learning_rate": 2e-05, + "loss": 0.0597, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.18564768135547638, + "learning_rate": 2e-05, + "loss": 0.0293, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 5.372669219970703, + "learning_rate": 2e-05, + "loss": 0.5941, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3632451593875885, + "learning_rate": 2e-05, + "loss": 0.1436, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.896437168121338, + "learning_rate": 2e-05, + "loss": 0.4735, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.8781577348709106, + "learning_rate": 2e-05, + "loss": 0.145, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.7643530368804932, + "learning_rate": 2e-05, + "loss": 0.6002, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 5.261754512786865, + "learning_rate": 2e-05, + "loss": 0.3395, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.46639588475227356, + "learning_rate": 2e-05, + "loss": 0.2251, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.405858039855957, + "learning_rate": 2e-05, + "loss": 0.2137, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.60910964012146, + "learning_rate": 2e-05, + "loss": 0.0769, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.4234938621520996, + "learning_rate": 2e-05, + "loss": 0.2421, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.10465911775827408, + "learning_rate": 2e-05, + "loss": 0.0088, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.4627697467803955, + "learning_rate": 2e-05, + "loss": 0.2024, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.365079641342163, + "learning_rate": 2e-05, + "loss": 0.0754, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.5433422327041626, + "learning_rate": 2e-05, + "loss": 0.0837, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.22659257054328918, + "learning_rate": 2e-05, + "loss": 0.0595, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.9471607208251953, + "learning_rate": 2e-05, + "loss": 0.0269, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.3210861682891846, + "learning_rate": 2e-05, + "loss": 0.0587, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 7.01611852645874, + "learning_rate": 2e-05, + "loss": 0.4037, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.2711181640625, + "learning_rate": 2e-05, + "loss": 0.2413, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.1692959070205688, + "learning_rate": 2e-05, + "loss": 0.0712, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.572306215763092, + "learning_rate": 2e-05, + "loss": 0.0201, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.557227611541748, + "learning_rate": 2e-05, + "loss": 0.5597, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.049049139022827, + "learning_rate": 2e-05, + "loss": 0.0954, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 10.485381126403809, + "learning_rate": 2e-05, + "loss": 1.5555, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.3989951610565186, + "learning_rate": 2e-05, + "loss": 0.2052, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 9.381972312927246, + "learning_rate": 2e-05, + "loss": 0.4917, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 5.929847717285156, + "learning_rate": 2e-05, + "loss": 0.6909, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.06029454618692398, + "learning_rate": 2e-05, + "loss": 0.0029, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.0042898654937744, + "learning_rate": 2e-05, + "loss": 0.0209, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.233729749917984, + "learning_rate": 2e-05, + "loss": 0.0198, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 5.546358585357666, + "learning_rate": 2e-05, + "loss": 0.4897, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.4404905140399933, + "learning_rate": 2e-05, + "loss": 0.2048, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.2590208053588867, + "learning_rate": 2e-05, + "loss": 0.0285, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.364134311676025, + "learning_rate": 2e-05, + "loss": 0.2614, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.654426097869873, + "learning_rate": 2e-05, + "loss": 0.3169, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.5928921699523926, + "learning_rate": 2e-05, + "loss": 0.1191, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 7.839958667755127, + "learning_rate": 2e-05, + "loss": 0.7773, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 5.929296970367432, + "learning_rate": 2e-05, + "loss": 0.13, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.7109313011169434, + "learning_rate": 2e-05, + "loss": 0.1082, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.560908555984497, + "learning_rate": 2e-05, + "loss": 0.1546, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206762417520640.0, + "train_loss": 0.2575809478759766, + "train_runtime": 59.7922, + "train_samples_per_second": 6.69, + "train_steps_per_second": 1.672 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206762417520640.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..c70164a13b60d394829728b72d106419c5805250 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:828f15d3077e7b33ca839c6a8d08473caf0c19d9ad881a4d15f2d06b6328418d +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..7e6f1a6b22b393e300444fd84c2f4ef7af8f95a5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d4745a3ff1c5de7f2435bf66d4dd723501a920229da515778f9e4679fc22f75 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..2178909f267e208259c36ef99738042cf0be3a8e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:673883fbea544de050a40a362f38d94ae4a40cb9d0f8fa1b910161b82925705d +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e99d36953aa6150e3e4889fde5abb00999f6917e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9ee0cac32efd5981e8c44b8fe10d8638f9fd6cea54641c1d3ae65ee150c62382 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..840a3dea02b7f6cc7a09be173aa72ac174fba7fc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:78011568f94c3f810fe8281ee5bc534a1d1b17cc26dcd4a2664bc293fbd5c011 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..929a4d79cdd87553cbe1641e84360cb1da079218 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:459646c1b9c949001e5de2ca426a6b724494eb2aa8c0b101878d8cd6cacdc2bd +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..1543efd9781fa46953f1a3e1938478af7756dbd7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d75190180d95c1dd056197f502582e8e1b5461f2328486bb9776a1c4ed7a7dd5 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8e4e6d3ecd26e87bae58df4022a6487fd47cce40 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5910e95fae4d793d23b4476425c1572a33f53e5c48a0a35470e6c4a41e43c16d +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..36893831549d04493a874a5a3d8792efba196952 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/17_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.191030740737915, + "learning_rate": 2e-05, + "loss": 0.0572, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.9470512866973877, + "learning_rate": 2e-05, + "loss": 0.3623, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.401775598526001, + "learning_rate": 2e-05, + "loss": 0.1392, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.3162297010421753, + "learning_rate": 2e-05, + "loss": 0.0695, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.324693202972412, + "learning_rate": 2e-05, + "loss": 0.1666, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.489325761795044, + "learning_rate": 2e-05, + "loss": 0.2007, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.145145297050476, + "learning_rate": 2e-05, + "loss": 0.037, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 6.8358378410339355, + "learning_rate": 2e-05, + "loss": 0.221, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 8.084610939025879, + "learning_rate": 2e-05, + "loss": 0.3403, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 7.938076496124268, + "learning_rate": 2e-05, + "loss": 0.9082, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.3662984371185303, + "learning_rate": 2e-05, + "loss": 0.1724, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 7.671472549438477, + "learning_rate": 2e-05, + "loss": 0.7248, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.355279564857483, + "learning_rate": 2e-05, + "loss": 0.1129, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.10215672850608826, + "learning_rate": 2e-05, + "loss": 0.0076, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.4111599922180176, + "learning_rate": 2e-05, + "loss": 0.1953, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.7624473571777344, + "learning_rate": 2e-05, + "loss": 0.1011, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.4505634307861328, + "learning_rate": 2e-05, + "loss": 0.0714, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.6107650399208069, + "learning_rate": 2e-05, + "loss": 0.0249, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.240485191345215, + "learning_rate": 2e-05, + "loss": 0.2469, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.658082365989685, + "learning_rate": 2e-05, + "loss": 0.2795, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.7464996576309204, + "learning_rate": 2e-05, + "loss": 0.0309, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.184717893600464, + "learning_rate": 2e-05, + "loss": 0.3748, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.7285668849945068, + "learning_rate": 2e-05, + "loss": 0.0345, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 4.990479946136475, + "learning_rate": 2e-05, + "loss": 0.2506, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.9914278984069824, + "learning_rate": 2e-05, + "loss": 0.3599, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.6742100715637207, + "learning_rate": 2e-05, + "loss": 0.2098, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.386765241622925, + "learning_rate": 2e-05, + "loss": 0.2318, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 5.5183892250061035, + "learning_rate": 2e-05, + "loss": 0.2923, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 7.482394218444824, + "learning_rate": 2e-05, + "loss": 0.4699, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.51076078414917, + "learning_rate": 2e-05, + "loss": 0.4798, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 9.347711563110352, + "learning_rate": 2e-05, + "loss": 1.1196, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 5.425720691680908, + "learning_rate": 2e-05, + "loss": 0.8726, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.46352750062942505, + "learning_rate": 2e-05, + "loss": 0.0666, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 7.798532009124756, + "learning_rate": 2e-05, + "loss": 0.325, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.46764373779296875, + "learning_rate": 2e-05, + "loss": 0.0313, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 4.014952659606934, + "learning_rate": 2e-05, + "loss": 0.4355, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.158535957336426, + "learning_rate": 2e-05, + "loss": 0.4122, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.0015740394592285, + "learning_rate": 2e-05, + "loss": 0.6463, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 5.4198689460754395, + "learning_rate": 2e-05, + "loss": 0.2875, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.5463374257087708, + "learning_rate": 2e-05, + "loss": 0.1226, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.9059522151947021, + "learning_rate": 2e-05, + "loss": 0.151, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.5052168369293213, + "learning_rate": 2e-05, + "loss": 0.3079, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 3.991971015930176, + "learning_rate": 2e-05, + "loss": 0.3958, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0204180479049683, + "learning_rate": 2e-05, + "loss": 0.0308, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 5.338406085968018, + "learning_rate": 2e-05, + "loss": 0.4781, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.4659030437469482, + "learning_rate": 2e-05, + "loss": 0.159, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.08972909301519394, + "learning_rate": 2e-05, + "loss": 0.0282, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.31429123878479, + "learning_rate": 2e-05, + "loss": 0.0678, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.024064060300588608, + "learning_rate": 2e-05, + "loss": 0.068, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.3935140371322632, + "learning_rate": 2e-05, + "loss": 0.0661, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2211830323740672.0, + "train_loss": 0.2648980140686035, + "train_runtime": 62.5857, + "train_samples_per_second": 6.391, + "train_steps_per_second": 1.598 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2211830323740672.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f35fbe5c858336e1adf252d97a05b58bd1120b3c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:755e6ef38b48afe07bd4feb6f436359c4185575888cb50d6d0324dd3ba9f7884 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..34e519a7eedbb083dcea69fd85148118bcbcfd3b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6c642fdc3c3153a3964235b3b9600fbae289506b70bc65246ac05c909fcb7c40 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..696080d42318e3bda8c948f665c62dbce908af8d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:270730e64b2d5672e48c2b9a4b3ea8fa80e63e4a24c12b257761687cc8ca3e20 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..c8a96de3aa43802518e0dd839aeca7c33df89fe4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b23975d6cb029c22eb3e456beb2267554595ff206a844d08345e17e6d2f3e26 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..f5bbf040b019b39d611f400e2d86ee47fcc38e85 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:66d30138b0b41aae59ff247d18efe3ef9486725f715fb102a063e7e8a8be4e14 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..3cbc0f3b55846babe184275440d73d16c1f2ecc9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ce79452d96c16cfcca475b09e60ddaa147c1cfee4d519a813f69f1884cced3c1 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..9d2995e87c009243f5b34a45d5e64228037ada41 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2291450ddaf1feafeddd06546f06a47ff44d9e531585b57d3bc56311cfc755fc +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..0a857edfc037803a96454eeb7010c28b6a202183 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2af296484e78f420795b21a42c961610f64e5255d6c898ec117b61baa3a4b93d +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..5610b39f864184d279d6c52db1797070ce639844 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/18_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.15653914213180542, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.6163437366485596, + "learning_rate": 2e-05, + "loss": 0.065, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.7129985690116882, + "learning_rate": 2e-05, + "loss": 0.0169, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.04419565945863724, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.057472262531518936, + "learning_rate": 2e-05, + "loss": 0.389, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 6.234667778015137, + "learning_rate": 2e-05, + "loss": 0.1637, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.611266851425171, + "learning_rate": 2e-05, + "loss": 0.0994, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.6677066087722778, + "learning_rate": 2e-05, + "loss": 0.0608, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.233916997909546, + "learning_rate": 2e-05, + "loss": 0.0377, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.5007081031799316, + "learning_rate": 2e-05, + "loss": 0.1398, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.08185774087905884, + "learning_rate": 2e-05, + "loss": 0.1353, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3986891806125641, + "learning_rate": 2e-05, + "loss": 0.075, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.6559237241744995, + "learning_rate": 2e-05, + "loss": 0.0739, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 5.084235191345215, + "learning_rate": 2e-05, + "loss": 0.5493, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.12226717174053192, + "learning_rate": 2e-05, + "loss": 0.0955, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 6.586478233337402, + "learning_rate": 2e-05, + "loss": 0.231, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.7441022396087646, + "learning_rate": 2e-05, + "loss": 0.0383, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.2272951900959015, + "learning_rate": 2e-05, + "loss": 0.1694, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.1148619651794434, + "learning_rate": 2e-05, + "loss": 0.0969, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.2628611922264099, + "learning_rate": 2e-05, + "loss": 0.0268, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.2279093265533447, + "learning_rate": 2e-05, + "loss": 0.3282, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.1913609504699707, + "learning_rate": 2e-05, + "loss": 0.0071, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.7044051885604858, + "learning_rate": 2e-05, + "loss": 0.3124, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.13698410987854, + "learning_rate": 2e-05, + "loss": 0.1211, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.3625404834747314, + "learning_rate": 2e-05, + "loss": 0.0811, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.03372323513031, + "learning_rate": 2e-05, + "loss": 0.1165, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 4.931954383850098, + "learning_rate": 2e-05, + "loss": 0.2928, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4350416958332062, + "learning_rate": 2e-05, + "loss": 0.3105, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.2746291756629944, + "learning_rate": 2e-05, + "loss": 0.0911, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8618091940879822, + "learning_rate": 2e-05, + "loss": 0.0765, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.01458148192614317, + "learning_rate": 2e-05, + "loss": 0.0364, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.9412890672683716, + "learning_rate": 2e-05, + "loss": 0.064, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.020536379888653755, + "learning_rate": 2e-05, + "loss": 0.1568, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.41321444511413574, + "learning_rate": 2e-05, + "loss": 0.0785, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 6.683852195739746, + "learning_rate": 2e-05, + "loss": 0.418, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.07003606855869293, + "learning_rate": 2e-05, + "loss": 0.0061, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.433330774307251, + "learning_rate": 2e-05, + "loss": 0.0765, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 9.671573638916016, + "learning_rate": 2e-05, + "loss": 0.3178, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 5.487994194030762, + "learning_rate": 2e-05, + "loss": 0.2352, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.4841979146003723, + "learning_rate": 2e-05, + "loss": 0.0127, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 5.6352643966674805, + "learning_rate": 2e-05, + "loss": 0.298, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.1305418759584427, + "learning_rate": 2e-05, + "loss": 0.1533, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.9235455989837646, + "learning_rate": 2e-05, + "loss": 0.0366, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.07865560054779053, + "learning_rate": 2e-05, + "loss": 0.0126, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.6884655952453613, + "learning_rate": 2e-05, + "loss": 0.1808, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.485936403274536, + "learning_rate": 2e-05, + "loss": 0.0909, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 6.493779182434082, + "learning_rate": 2e-05, + "loss": 0.297, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.5242996215820312, + "learning_rate": 2e-05, + "loss": 0.0546, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.030198778957128525, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 6.102869033813477, + "learning_rate": 2e-05, + "loss": 0.2856, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206568850391040.0, + "train_loss": 0.1403981304168701, + "train_runtime": 61.9338, + "train_samples_per_second": 6.459, + "train_steps_per_second": 1.615 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206568850391040.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..b9f6d1fb50ea69d29d5ebe2224b33af955c7b924 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f515f54dc5dc065ea6b851a0eb13091134335512129e4bb1368db48fad3c868e +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..c0aaa69bb99b2213a671c07043478cdc4c8c2775 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cd3e277c57f423b114bc947c7f49e7394a6e98217c1c339b88663a4d2f453b39 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..4d032ec199e3563f07bbc5693ab941185b4cd9f6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3648157134af2f32a21326a4b7927851f160244abd3107781382c17e94d44ffa +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d36f841e4fa3bc49e6ac46403c166a852e1873c8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd45f2b785cd464b57b1828d234deaf4f5610140bfe896d417b76b061d5a98b1 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..15065d4991ef2890fb9fcef9f01cc9249df7598b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cf5242349e3b30c5f1b0211488bf1d496022ca43aa06a9d08874fe0d4dd3968d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..dcd2f5923934d5cd53c5919944cb8f8fefb2b611 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23b48347f8c39dd717fdfd9341cb008a574b8e9c1a56be810d7265a3a5ba640b +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b3e92200880cfd169ae6f49e525ab492f2fad309 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f54e14143861c951d7bd71a69ca51cd32be4d5c1f13b0a33a4875bd6ca40bf06 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..4df46bf9f1b3fa0447dd5d1ad9f1741ff716fd06 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4f0876e5e1bfe0191345d22ae5b1d389b5bc30bdcf0e08829cc88d63d723b50f +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..67c7a37d305c32aef33dbd13d6b5132724a6aaba --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/19_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.4283304214477539, + "learning_rate": 2e-05, + "loss": 0.0895, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.4142312705516815, + "learning_rate": 2e-05, + "loss": 0.0922, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.640970766544342, + "learning_rate": 2e-05, + "loss": 0.0571, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.0663375854492188, + "learning_rate": 2e-05, + "loss": 0.1486, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.782832622528076, + "learning_rate": 2e-05, + "loss": 0.3524, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.5954317450523376, + "learning_rate": 2e-05, + "loss": 0.0481, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.4493829011917114, + "learning_rate": 2e-05, + "loss": 0.0798, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.7121614813804626, + "learning_rate": 2e-05, + "loss": 0.1457, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.22501978278160095, + "learning_rate": 2e-05, + "loss": 0.2599, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.20836865901947021, + "learning_rate": 2e-05, + "loss": 0.0139, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.0013967752456665, + "learning_rate": 2e-05, + "loss": 0.0781, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.2661224603652954, + "learning_rate": 2e-05, + "loss": 0.0482, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.835895538330078, + "learning_rate": 2e-05, + "loss": 0.1743, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.8032569885253906, + "learning_rate": 2e-05, + "loss": 0.1341, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.0321309566497803, + "learning_rate": 2e-05, + "loss": 0.2195, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.4017966091632843, + "learning_rate": 2e-05, + "loss": 0.4796, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.166468858718872, + "learning_rate": 2e-05, + "loss": 0.2156, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.34434056282043457, + "learning_rate": 2e-05, + "loss": 0.0939, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2946701645851135, + "learning_rate": 2e-05, + "loss": 0.0366, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.10451260209083557, + "learning_rate": 2e-05, + "loss": 0.0623, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.337615728378296, + "learning_rate": 2e-05, + "loss": 0.0633, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7607973217964172, + "learning_rate": 2e-05, + "loss": 0.5788, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.9001750349998474, + "learning_rate": 2e-05, + "loss": 0.1544, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.2669612467288971, + "learning_rate": 2e-05, + "loss": 0.2939, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.34477338194847107, + "learning_rate": 2e-05, + "loss": 0.2766, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.125678062438965, + "learning_rate": 2e-05, + "loss": 0.2522, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.3169984221458435, + "learning_rate": 2e-05, + "loss": 0.1663, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.0864514485001564, + "learning_rate": 2e-05, + "loss": 0.0136, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.9170241355895996, + "learning_rate": 2e-05, + "loss": 0.1704, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.7182433605194092, + "learning_rate": 2e-05, + "loss": 0.0799, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.11117897927761078, + "learning_rate": 2e-05, + "loss": 0.178, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.636600494384766, + "learning_rate": 2e-05, + "loss": 0.2864, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.297088146209717, + "learning_rate": 2e-05, + "loss": 0.2748, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.049335215240716934, + "learning_rate": 2e-05, + "loss": 0.083, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.4287002086639404, + "learning_rate": 2e-05, + "loss": 0.4071, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.4949915111064911, + "learning_rate": 2e-05, + "loss": 0.1751, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.24649862945079803, + "learning_rate": 2e-05, + "loss": 0.0265, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 7.960385322570801, + "learning_rate": 2e-05, + "loss": 0.5993, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.927978515625, + "learning_rate": 2e-05, + "loss": 0.1427, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.474682092666626, + "learning_rate": 2e-05, + "loss": 0.1558, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.548583507537842, + "learning_rate": 2e-05, + "loss": 0.3386, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.1477077305316925, + "learning_rate": 2e-05, + "loss": 0.02, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0029542515985667706, + "learning_rate": 2e-05, + "loss": 0.0492, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.1292826384305954, + "learning_rate": 2e-05, + "loss": 0.3385, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.7265453338623047, + "learning_rate": 2e-05, + "loss": 0.0571, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.9342368841171265, + "learning_rate": 2e-05, + "loss": 0.0591, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.5672251582145691, + "learning_rate": 2e-05, + "loss": 0.09, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.283287525177002, + "learning_rate": 2e-05, + "loss": 0.1767, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.19825570285320282, + "learning_rate": 2e-05, + "loss": 0.0373, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.07856085151433945, + "learning_rate": 2e-05, + "loss": 0.3329, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293519879012352.0, + "train_loss": 0.1741359758377075, + "train_runtime": 104.464, + "train_samples_per_second": 3.829, + "train_steps_per_second": 0.957 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293519879012352.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..2ff0fda8db101e292c7b2a52e05989c78c9ce9c0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ae9e33458a5800c851329c83b5e2fbb4200e552638f1609afdece9bda2d051c5 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..318c7526708387fdedfd9d10f6c4970ab1508da6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:649e315fa3f28f5215f32b6d3bf16062e3c17c3d6c345bea8ff6c67d540fef23 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..b6d4fc4207645de3258682c7591337a68a311830 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:48bb99680df7c74985c8141de69f5c6d4b2f9aaaf6542e94d295e5624e0b8948 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..8b57c6b84511ababb5047af9cc7625a3574d98c5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d432ba5aa5ce12cc14322f0bb4d55a5305c04bedd5104c5ed92823349dd321c8 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..d054eb461044644b129c829721584647ca21b1f7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0b5362175b75b1c5a88e686bc8f80bf7e162e54fd0171c8d630eb6b1184719e4 +size 368442474 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..f35f23032dbbe6179173aad755904b624cf339df --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9a7edcad69f767e0517d1b00e46df60efeecd64f483bca51aaa7984164b16da6 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..786ad52fe459420f63a287fcd430920f0e90e82c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:565b1485c8da544c001e3897abc4cdc1f5aec37e44044056ee97adca8f5aadb0 +size 368442474 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..b71955d0029a758a12b9ed64eb45ba60eb02731b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:26c5510390fb6b3bdb9d0888e880fbde8177037c9d14aede98f99e008deed70e +size 368442474 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..0fcc670c84ce3a75dbfff695d937f60fa6469f68 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/1_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.04779736325144768, + "learning_rate": 2e-05, + "loss": 0.0054, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.04407791793346405, + "learning_rate": 2e-05, + "loss": 0.0266, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.10724038630723953, + "learning_rate": 2e-05, + "loss": 0.003, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.22232206165790558, + "learning_rate": 2e-05, + "loss": 0.0081, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.1078585684299469, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.1512720584869385, + "learning_rate": 2e-05, + "loss": 0.0362, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.0208271536976099, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.19057884812355042, + "learning_rate": 2e-05, + "loss": 0.055, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.07116301357746124, + "learning_rate": 2e-05, + "loss": 0.06, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.006559637375175953, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.037236738950014114, + "learning_rate": 2e-05, + "loss": 0.0158, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.05198603495955467, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.7721407413482666, + "learning_rate": 2e-05, + "loss": 0.1239, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.02655797451734543, + "learning_rate": 2e-05, + "loss": 0.0009, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.24767422676086426, + "learning_rate": 2e-05, + "loss": 0.0054, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.19544078409671783, + "learning_rate": 2e-05, + "loss": 0.0095, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.0685712099075317, + "learning_rate": 2e-05, + "loss": 0.0204, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.002792870858684182, + "learning_rate": 2e-05, + "loss": 0.0002, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.06202378869056702, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.066880702972412, + "learning_rate": 2e-05, + "loss": 0.0215, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.3063740134239197, + "learning_rate": 2e-05, + "loss": 0.0052, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.760720729827881, + "learning_rate": 2e-05, + "loss": 0.1908, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.004246275406330824, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.012776656076312065, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.0069309682585299015, + "learning_rate": 2e-05, + "loss": 0.0009, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.04661623388528824, + "learning_rate": 2e-05, + "loss": 0.033, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.912585496902466, + "learning_rate": 2e-05, + "loss": 0.073, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.04052748158574104, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.013371194712817669, + "learning_rate": 2e-05, + "loss": 0.1238, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.005390165839344263, + "learning_rate": 2e-05, + "loss": 0.0255, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.243587464094162, + "learning_rate": 2e-05, + "loss": 0.0037, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.017862478271126747, + "learning_rate": 2e-05, + "loss": 0.0087, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.014391588047146797, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.19254015386104584, + "learning_rate": 2e-05, + "loss": 0.0045, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.009134696796536446, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.8035869002342224, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.7504981160163879, + "learning_rate": 2e-05, + "loss": 0.0164, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.013966098427772522, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.649307370185852, + "learning_rate": 2e-05, + "loss": 0.0325, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.00970319751650095, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.02150818333029747, + "learning_rate": 2e-05, + "loss": 0.0054, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0037529587280005217, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.017755217850208282, + "learning_rate": 2e-05, + "loss": 0.1134, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.03316999226808548, + "learning_rate": 2e-05, + "loss": 0.0015, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.15144549310207367, + "learning_rate": 2e-05, + "loss": 0.0033, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.13540761172771454, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.016004299744963646, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.4086146056652069, + "learning_rate": 2e-05, + "loss": 0.0067, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.007050918880850077, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.031060179695487022, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2217003838341120.0, + "train_loss": 0.02136661410331726, + "train_runtime": 64.3651, + "train_samples_per_second": 6.215, + "train_steps_per_second": 1.554 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2217003838341120.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..61827803aa0566b7d4ebb99ad981847b05874f5a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:15c6634d632e56eefbac7b0d7be2e8bc741011dd0aa61bb564d67d0ca399537d +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ac5be816f7188fbd3892d87eaccb758796681cf8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f9bfa955e5a7751a9a45d51e371107a0135f92d7830a7dfc44785ce3d5e9ad2e +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..88b74c1e76ae5b660347a1a2020c37c425cce288 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1082b16769329ea2b1960881a1840873d096a7fa8a516723dcddd889c7e3583e +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..81c11b4723d3a222a5b960789a9f64c4a7cde7fa --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3c7c806bb2d6f0c1dc29927bfbb770f5550191c724ce856ac13274244d4d857b +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..c7f5559f6e3896529a79316cb977886e23e4d41a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6f6d268d27975eee69ab7648c1c361bab4ad4abe4ddc28ba5e7e048d34575290 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..f49c778814be50dc3226852ca5e8cb13cbd54040 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2aa1b77f61834fe7235f38789d537d087daae639730233670cc7849d8b2ff09 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..8693399a2ca2d4b39bab37ac680318db4099966c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:169cc1d3430c04fe556489cdde30b7bf7f9c3389265de668ce7de532103397ad +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..5245cb3503f069eee6bea63e72953be2f6d02dea --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bbf1e01cd865c3c5f92b5b96e2e79c4d3fa7c2662e3d36c88e91eeded1ae7154 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..ed2f9dccd0ca677ad93a0c7c65a5d0fd922be946 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/20_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.12505054473876953, + "learning_rate": 2e-05, + "loss": 0.0269, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.8414512872695923, + "learning_rate": 2e-05, + "loss": 0.0781, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 4.093201637268066, + "learning_rate": 2e-05, + "loss": 0.1459, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.7235778570175171, + "learning_rate": 2e-05, + "loss": 0.0244, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.993597507476807, + "learning_rate": 2e-05, + "loss": 0.5012, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.10781623423099518, + "learning_rate": 2e-05, + "loss": 0.0063, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.5678093433380127, + "learning_rate": 2e-05, + "loss": 0.0433, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.2410985231399536, + "learning_rate": 2e-05, + "loss": 0.0775, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.8636186122894287, + "learning_rate": 2e-05, + "loss": 0.186, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.07627388834953308, + "learning_rate": 2e-05, + "loss": 0.1664, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.5126129984855652, + "learning_rate": 2e-05, + "loss": 0.0159, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3146267235279083, + "learning_rate": 2e-05, + "loss": 0.0147, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.0028334856033325, + "learning_rate": 2e-05, + "loss": 0.2845, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.2397547960281372, + "learning_rate": 2e-05, + "loss": 0.0129, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.9763955473899841, + "learning_rate": 2e-05, + "loss": 0.0249, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.188786745071411, + "learning_rate": 2e-05, + "loss": 0.2365, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.11952368915081024, + "learning_rate": 2e-05, + "loss": 0.039, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.5312955379486084, + "learning_rate": 2e-05, + "loss": 0.0357, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.43194249272346497, + "learning_rate": 2e-05, + "loss": 0.0871, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.5910495519638062, + "learning_rate": 2e-05, + "loss": 0.0623, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.5835009217262268, + "learning_rate": 2e-05, + "loss": 0.1518, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.4090814590454102, + "learning_rate": 2e-05, + "loss": 0.0813, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.7141611576080322, + "learning_rate": 2e-05, + "loss": 0.1269, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.19910728931427, + "learning_rate": 2e-05, + "loss": 0.0338, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.4662742614746094, + "learning_rate": 2e-05, + "loss": 0.1593, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.2152310609817505, + "learning_rate": 2e-05, + "loss": 0.0085, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.4730389714241028, + "learning_rate": 2e-05, + "loss": 0.2097, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.0842965841293335, + "learning_rate": 2e-05, + "loss": 0.035, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.013021734543144703, + "learning_rate": 2e-05, + "loss": 0.2029, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.12087839841842651, + "learning_rate": 2e-05, + "loss": 0.0537, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.104187250137329, + "learning_rate": 2e-05, + "loss": 0.2469, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.9862236976623535, + "learning_rate": 2e-05, + "loss": 0.6018, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 6.924476146697998, + "learning_rate": 2e-05, + "loss": 0.8462, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.247273921966553, + "learning_rate": 2e-05, + "loss": 0.1949, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.757524311542511, + "learning_rate": 2e-05, + "loss": 0.0454, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.08681836724281311, + "learning_rate": 2e-05, + "loss": 0.0357, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.18200421333312988, + "learning_rate": 2e-05, + "loss": 0.0124, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 6.961614608764648, + "learning_rate": 2e-05, + "loss": 0.6555, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.643007755279541, + "learning_rate": 2e-05, + "loss": 0.03, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.019762586802244186, + "learning_rate": 2e-05, + "loss": 0.0184, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.15375632047653198, + "learning_rate": 2e-05, + "loss": 0.0812, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.9476636052131653, + "learning_rate": 2e-05, + "loss": 0.0953, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.3325204849243164, + "learning_rate": 2e-05, + "loss": 0.0217, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.6414198875427246, + "learning_rate": 2e-05, + "loss": 0.1609, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.5200161933898926, + "learning_rate": 2e-05, + "loss": 0.141, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.9570791721343994, + "learning_rate": 2e-05, + "loss": 0.2977, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.9023889899253845, + "learning_rate": 2e-05, + "loss": 0.0386, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.11767003685235977, + "learning_rate": 2e-05, + "loss": 0.0043, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.351616382598877, + "learning_rate": 2e-05, + "loss": 0.1316, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.2494819611310959, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293985824243712.0, + "train_loss": 0.1360475528240204, + "train_runtime": 93.7234, + "train_samples_per_second": 4.268, + "train_steps_per_second": 1.067 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293985824243712.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f8f92509170382fbc008c6edb42c27e2b0c5a6eb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56765ce8266a7227a1c94bac30c20bed27a9362b36de4381b9d48b45d3f1c120 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..6a84b073861ad4331b0c8fd8d61c9ef52b80e568 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f6a35d4c7ccf9d3b5ec24b1c186b3ab3c47564b5f2afffe8bbb5e9f0300a7fa9 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..3105a418fcd4426cab3b467208a5de3c1d4cb742 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ae586130bbaa1964bf2f3509894c0ccb456084a694bd1666a3d39ac146de781b +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..808ab5297e412f6c717ac06302020dd5b09924d0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6484468600ae2da5a40522980cb0b839e1fc2cf5b924678ac5ce59a926e87ba7 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..be07652189dabbafa2d075aa0567b49c8333d317 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:021414698e067a8b7e51da2632cdbc9770e80fbc9a48ac26baaa6ad15d6ac227 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..17c66d2b80d472a4b3b579ca85db0c3fc5f56084 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:79b626c816e1a571cdb1a6248d647a3a09459f65eedf67b3587042a0f8fe1f81 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..e0077be807ebffbd7fa79722f44fa4a4c2f61227 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e2e88f8ccf5030be4ca6f68b012e4989d322e79cbe6a2a50353a93794299da3c +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..0805ee80fc6c33b74507d9888c0742cec7a0179f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:36178a35922c2a7023258b4ecaab56579665e3656f75fdaefecaab548c4a4dbc +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..a00a8829190c018f775a5cf23b5dd45a7deb5f29 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/2_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.0204379558563232, + "learning_rate": 2e-05, + "loss": 0.1828, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.346928596496582, + "learning_rate": 2e-05, + "loss": 0.7148, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.4880144596099854, + "learning_rate": 2e-05, + "loss": 0.3557, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.5174741744995117, + "learning_rate": 2e-05, + "loss": 0.3568, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.0307538509368896, + "learning_rate": 2e-05, + "loss": 0.1474, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.730238199234009, + "learning_rate": 2e-05, + "loss": 0.497, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.3734097480773926, + "learning_rate": 2e-05, + "loss": 0.1608, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.6674742698669434, + "learning_rate": 2e-05, + "loss": 0.3125, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.927203893661499, + "learning_rate": 2e-05, + "loss": 0.2602, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.0440961122512817, + "learning_rate": 2e-05, + "loss": 0.4749, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.406931608915329, + "learning_rate": 2e-05, + "loss": 0.0753, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.9952826499938965, + "learning_rate": 2e-05, + "loss": 0.2558, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.368868350982666, + "learning_rate": 2e-05, + "loss": 0.5076, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.997451066970825, + "learning_rate": 2e-05, + "loss": 0.1689, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.41830164194107056, + "learning_rate": 2e-05, + "loss": 0.0824, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.0760037899017334, + "learning_rate": 2e-05, + "loss": 0.1892, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.387637138366699, + "learning_rate": 2e-05, + "loss": 0.3162, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.9407545328140259, + "learning_rate": 2e-05, + "loss": 0.0502, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.0937774181365967, + "learning_rate": 2e-05, + "loss": 0.5708, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 5.482045650482178, + "learning_rate": 2e-05, + "loss": 0.4185, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.1550674438476562, + "learning_rate": 2e-05, + "loss": 0.3374, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.635411262512207, + "learning_rate": 2e-05, + "loss": 0.0841, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.9332603216171265, + "learning_rate": 2e-05, + "loss": 0.2162, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.3021145164966583, + "learning_rate": 2e-05, + "loss": 0.1273, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.3171757459640503, + "learning_rate": 2e-05, + "loss": 0.1258, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.9088447093963623, + "learning_rate": 2e-05, + "loss": 0.3096, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.889411211013794, + "learning_rate": 2e-05, + "loss": 0.2849, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 4.403746604919434, + "learning_rate": 2e-05, + "loss": 0.2646, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.4958213567733765, + "learning_rate": 2e-05, + "loss": 0.0634, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.14780236780643463, + "learning_rate": 2e-05, + "loss": 0.1105, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.0633269548416138, + "learning_rate": 2e-05, + "loss": 0.3269, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.0896058082580566, + "learning_rate": 2e-05, + "loss": 0.1519, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.7947436571121216, + "learning_rate": 2e-05, + "loss": 0.1026, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.9582667350769043, + "learning_rate": 2e-05, + "loss": 0.2481, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.6318299770355225, + "learning_rate": 2e-05, + "loss": 0.3735, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.2842870950698853, + "learning_rate": 2e-05, + "loss": 0.4752, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.9569968581199646, + "learning_rate": 2e-05, + "loss": 0.0334, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.0460015535354614, + "learning_rate": 2e-05, + "loss": 0.1551, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.8575685620307922, + "learning_rate": 2e-05, + "loss": 0.0423, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 6.298202991485596, + "learning_rate": 2e-05, + "loss": 0.5475, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 8.088481903076172, + "learning_rate": 2e-05, + "loss": 1.7529, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.19255176186561584, + "learning_rate": 2e-05, + "loss": 0.0097, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.6745517253875732, + "learning_rate": 2e-05, + "loss": 0.2633, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 4.204707622528076, + "learning_rate": 2e-05, + "loss": 0.3396, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.493457317352295, + "learning_rate": 2e-05, + "loss": 0.9067, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.209862619638443, + "learning_rate": 2e-05, + "loss": 0.1157, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.1772608757019043, + "learning_rate": 2e-05, + "loss": 0.2423, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.7324260473251343, + "learning_rate": 2e-05, + "loss": 0.1461, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.583072662353516, + "learning_rate": 2e-05, + "loss": 0.6593, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.9615267515182495, + "learning_rate": 2e-05, + "loss": 0.0971, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5221258941693952.0, + "train_loss": 0.300220468044281, + "train_runtime": 100.882, + "train_samples_per_second": 3.965, + "train_steps_per_second": 0.991 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5221258941693952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..703d7eca71ff6a345ec6ffd345e277816e6355d9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:beb6b05fca1b9ec8c28723f8c56f9863d3c64d33f408aaa4567d34372ae726a6 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..bb7381a153d081c606eb611e4fb1be2fc5a470f2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:404911f228ecafa4a5d68db6f3104dc321c2bae644a1d5ff8bec2838e1e2b40e +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a31991cfacee684e1fa085be193b7091276f58b5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7601caeff49bc05f1b62e80866bfda769e72d8eaa0a734a2d91c92130970a0c0 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..97fb8597f17c75688530044bdc477197c13d3d09 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7a45f7ebdf2dce7f052431f53478da3f60eda2379545a52f93471cd7486425e2 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..2811d3fe4c46a7660807fc6c1a77615b5fb44b11 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:358803b003062b748192a176db42930f5b9bc35e8c76278ef4ea37f9dedbad94 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..90a350daf2656817d0a54a18e40e95bc3b5795a2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2cf76291203b038b9913d60a6fa0cf960820a21aa9c96fd3f66ed4e488d71c5e +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..cc1c99c3a04cde5c5007577e56bf75216af01f55 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c49a15df090b3247b64a6020baeaf00a606f007cf8c98dc4a3acd4aef7045d77 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d92b9efcea7c2b28398c54034b09ee4b691b7081 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fe21c6b079e3ccf0a208b7e3daf460448223132277cc5aec1e01b3d5299b0b3e +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c4a7104c7c3c914f285a752a3ad58387eab016a4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/3_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.5754480361938477, + "learning_rate": 2e-05, + "loss": 0.7847, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.3294107913970947, + "learning_rate": 2e-05, + "loss": 0.3578, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.232266426086426, + "learning_rate": 2e-05, + "loss": 0.2612, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.900785446166992, + "learning_rate": 2e-05, + "loss": 0.7296, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.120523929595947, + "learning_rate": 2e-05, + "loss": 0.5941, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 4.678377151489258, + "learning_rate": 2e-05, + "loss": 0.785, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.6316890716552734, + "learning_rate": 2e-05, + "loss": 0.7632, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.762193202972412, + "learning_rate": 2e-05, + "loss": 0.6228, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.2454519271850586, + "learning_rate": 2e-05, + "loss": 0.1599, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.3454760313034058, + "learning_rate": 2e-05, + "loss": 0.3537, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.5215377807617188, + "learning_rate": 2e-05, + "loss": 0.2427, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.9892505407333374, + "learning_rate": 2e-05, + "loss": 0.5273, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.4746805429458618, + "learning_rate": 2e-05, + "loss": 0.201, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 4.057143688201904, + "learning_rate": 2e-05, + "loss": 0.3147, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.1324379444122314, + "learning_rate": 2e-05, + "loss": 0.3932, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.4833552837371826, + "learning_rate": 2e-05, + "loss": 0.3722, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.6996318101882935, + "learning_rate": 2e-05, + "loss": 0.6477, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.521028518676758, + "learning_rate": 2e-05, + "loss": 0.7432, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.1221779584884644, + "learning_rate": 2e-05, + "loss": 0.1562, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.5840959548950195, + "learning_rate": 2e-05, + "loss": 0.5464, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.5873329639434814, + "learning_rate": 2e-05, + "loss": 0.6084, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 6.273464202880859, + "learning_rate": 2e-05, + "loss": 0.6902, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.1911542415618896, + "learning_rate": 2e-05, + "loss": 0.2072, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.9990672469139099, + "learning_rate": 2e-05, + "loss": 0.0698, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 3.3292887210845947, + "learning_rate": 2e-05, + "loss": 0.5413, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.7526795864105225, + "learning_rate": 2e-05, + "loss": 0.6842, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.0994989201426506, + "learning_rate": 2e-05, + "loss": 0.1802, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.3864864110946655, + "learning_rate": 2e-05, + "loss": 0.2625, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.462929368019104, + "learning_rate": 2e-05, + "loss": 0.3265, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 3.693169593811035, + "learning_rate": 2e-05, + "loss": 0.5219, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.6277828216552734, + "learning_rate": 2e-05, + "loss": 0.3768, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.624203681945801, + "learning_rate": 2e-05, + "loss": 0.2764, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.491536617279053, + "learning_rate": 2e-05, + "loss": 0.5079, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 4.2294158935546875, + "learning_rate": 2e-05, + "loss": 0.7836, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.7548465728759766, + "learning_rate": 2e-05, + "loss": 0.4219, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.609562873840332, + "learning_rate": 2e-05, + "loss": 0.6909, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.6600353717803955, + "learning_rate": 2e-05, + "loss": 0.696, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 4.062559127807617, + "learning_rate": 2e-05, + "loss": 0.2884, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.9801759719848633, + "learning_rate": 2e-05, + "loss": 0.3433, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.7080295085906982, + "learning_rate": 2e-05, + "loss": 0.1612, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.8384615182876587, + "learning_rate": 2e-05, + "loss": 0.1051, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.708599805831909, + "learning_rate": 2e-05, + "loss": 0.2158, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.6118149757385254, + "learning_rate": 2e-05, + "loss": 0.5977, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0794622898101807, + "learning_rate": 2e-05, + "loss": 0.3154, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3016270399093628, + "learning_rate": 2e-05, + "loss": 0.5135, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.5667765140533447, + "learning_rate": 2e-05, + "loss": 0.3341, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.0138673782348633, + "learning_rate": 2e-05, + "loss": 0.2283, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.9124984741210938, + "learning_rate": 2e-05, + "loss": 0.2311, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.7780407667160034, + "learning_rate": 2e-05, + "loss": 0.3617, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.957095742225647, + "learning_rate": 2e-05, + "loss": 0.2057, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5408561442062336.0, + "train_loss": 0.4260712242126465, + "train_runtime": 94.6812, + "train_samples_per_second": 4.225, + "train_steps_per_second": 1.056 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5408561442062336.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5ebf7f441abe4d4cc22a8d84a5ad1d631a94690c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:91793a82d9574d7b3534e91946dcbb8164e5ce487e2688ad3b3b531148d4e436 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..2bed4e26624b27e96f9cd963a5a692f26e2b4a88 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:15dfed66f08ff0da91ae2a1765c84972a4e050ae3d332da395c6af81b8535600 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..d40b2611ebd726c072d5f44439914124e9a588a1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1e585a6dfc607530a708c1c2d27798293d1a3a424466ad57054af37b806cad7b +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..6fbc7c2b61d47e0c35cbcf03b98835180e35102c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92bb46186123a26aa9eb299f55aa37dfed0add13c84837fe89691191d228f351 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..93dc8c0b5d79019dcfd313e68f7969fa0cbc1876 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:044cef7a71453e2d447eb1e06e75907a0b185ae91adf66cff3473cd38b7f0f7b +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d433d41042820213dc014e2ff8df2471da72920b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be3b50ffadd8ca396fe664f80c20d052e7abc058dc5d29ac094cd4562170234d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..334a2dbda26ad918c202d16777605c22f1a0fd73 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0ce6606c9e7b8c4674fef937a9b52f23ebe7feba968ddce46747fb26218694f0 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..1c91801f34a6345a31b52b82c48957a1cf8edf94 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c1f463ca30f7145b24f49dd88d6c1d8659a1c60775d2b1a809f15d51c08eb0ea +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..49e0ffb8292fbe6428ba3c9913a2a82195798524 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/4_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.0798341035842896, + "learning_rate": 2e-05, + "loss": 0.2278, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.8705922365188599, + "learning_rate": 2e-05, + "loss": 0.4041, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.8320990800857544, + "learning_rate": 2e-05, + "loss": 0.3081, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.787644386291504, + "learning_rate": 2e-05, + "loss": 0.3848, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.0484519004821777, + "learning_rate": 2e-05, + "loss": 0.2087, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.7446677684783936, + "learning_rate": 2e-05, + "loss": 0.1328, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.147367000579834, + "learning_rate": 2e-05, + "loss": 0.333, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.211151123046875, + "learning_rate": 2e-05, + "loss": 0.5671, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.735642671585083, + "learning_rate": 2e-05, + "loss": 0.1593, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.2057065963745117, + "learning_rate": 2e-05, + "loss": 0.1962, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.6223304271698, + "learning_rate": 2e-05, + "loss": 1.0657, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.0393385887145996, + "learning_rate": 2e-05, + "loss": 0.5076, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.645646572113037, + "learning_rate": 2e-05, + "loss": 0.3244, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.828223705291748, + "learning_rate": 2e-05, + "loss": 0.4552, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.9316736459732056, + "learning_rate": 2e-05, + "loss": 0.3821, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.433586359024048, + "learning_rate": 2e-05, + "loss": 0.356, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.9435164928436279, + "learning_rate": 2e-05, + "loss": 0.214, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.403280019760132, + "learning_rate": 2e-05, + "loss": 0.5273, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.6494677066802979, + "learning_rate": 2e-05, + "loss": 0.2252, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.68591833114624, + "learning_rate": 2e-05, + "loss": 0.6289, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.4931756258010864, + "learning_rate": 2e-05, + "loss": 0.3936, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.7911590337753296, + "learning_rate": 2e-05, + "loss": 0.2837, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.626120090484619, + "learning_rate": 2e-05, + "loss": 0.375, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.796732783317566, + "learning_rate": 2e-05, + "loss": 0.3422, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.6290534734725952, + "learning_rate": 2e-05, + "loss": 0.2075, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.9719321727752686, + "learning_rate": 2e-05, + "loss": 0.7373, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.8254743814468384, + "learning_rate": 2e-05, + "loss": 0.2916, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.9815352559089661, + "learning_rate": 2e-05, + "loss": 0.1304, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.0491750240325928, + "learning_rate": 2e-05, + "loss": 0.3623, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.7500059604644775, + "learning_rate": 2e-05, + "loss": 0.4126, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.2983802556991577, + "learning_rate": 2e-05, + "loss": 0.2244, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.528508186340332, + "learning_rate": 2e-05, + "loss": 0.4946, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.6394054889678955, + "learning_rate": 2e-05, + "loss": 0.3512, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.6836693286895752, + "learning_rate": 2e-05, + "loss": 0.3, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 4.019546985626221, + "learning_rate": 2e-05, + "loss": 0.3889, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.8557461500167847, + "learning_rate": 2e-05, + "loss": 0.2444, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.8953967094421387, + "learning_rate": 2e-05, + "loss": 0.29, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.8404831886291504, + "learning_rate": 2e-05, + "loss": 0.5609, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.9191360473632812, + "learning_rate": 2e-05, + "loss": 0.3447, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.7631478309631348, + "learning_rate": 2e-05, + "loss": 0.2975, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.5603580474853516, + "learning_rate": 2e-05, + "loss": 0.2434, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.6434848308563232, + "learning_rate": 2e-05, + "loss": 0.3217, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.4387174844741821, + "learning_rate": 2e-05, + "loss": 0.155, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0281782150268555, + "learning_rate": 2e-05, + "loss": 0.2226, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.4844415187835693, + "learning_rate": 2e-05, + "loss": 0.394, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.2288904190063477, + "learning_rate": 2e-05, + "loss": 0.295, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.765199899673462, + "learning_rate": 2e-05, + "loss": 0.4413, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.4653961658477783, + "learning_rate": 2e-05, + "loss": 0.3553, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.230545163154602, + "learning_rate": 2e-05, + "loss": 0.0959, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 5.263822078704834, + "learning_rate": 2e-05, + "loss": 0.5595, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 6047770452426752.0, + "train_loss": 0.35449752807617185, + "train_runtime": 104.9113, + "train_samples_per_second": 3.813, + "train_steps_per_second": 0.953 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 6047770452426752.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..58163b3c78b08a7f68ea8ca9ba5be855897aefca --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b5f08618c7da45d2f4f16f23526cab6fc28c685e7df125b6e9ea81571be5e18d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..6cac01f749ccc41e03b48110a745ad5d5da9b988 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:30019189be8acf1f5216f843129d7c3dbad0767865b84308464be63a482ada50 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..bed53e0f25e766d14ea3b73dc83d6c3ef59b3308 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c9b220d9985a6334c3b4e60b4f8653078bf63c2e62c5d3873d5de688d7b860e +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..de21bde15b5838707d85fecb0fb8321c17821e95 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:55475c87b0d15c645054648a470a72bfd23d72c49e2fb9f5ca718a48a91bb4c7 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..ec3d4f83fec7f5e3bae55389414e28a3d391daa6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dbca3aabda6d6b754416a9a3b1fe4f1af9f19a62fdd975969de0b4edd18419fd +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..38ae5c5d5c28422c3e64d2f1170678a136d5cd58 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:929e4617bb38d645b4066fabf2af746d43012ffa6f6fa7c698c6fd49706dc533 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..5757dee77bc89fd22371ae4c1eb22ab317c1397a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e24d60c1e02eb739698fd6f7940b061ea274aab265db4a587e918e2ad3707fbd +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..35c88a8da1bee6a30dea0a49755858007a1f2a44 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6f37c4c07f0843cbad8f8a601230f212024ba9a806e820b2ac94e9205ce3300f +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..da326f85a4f205a322572d2f3590027f04df8d06 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/5_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.489640951156616, + "learning_rate": 2e-05, + "loss": 0.1587, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.661216974258423, + "learning_rate": 2e-05, + "loss": 0.3314, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.0923573970794678, + "learning_rate": 2e-05, + "loss": 0.1311, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.12899190187454224, + "learning_rate": 2e-05, + "loss": 0.1164, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.9320750832557678, + "learning_rate": 2e-05, + "loss": 0.1287, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.1474609225988388, + "learning_rate": 2e-05, + "loss": 0.0295, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.013700793497264385, + "learning_rate": 2e-05, + "loss": 0.0303, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.09249375760555267, + "learning_rate": 2e-05, + "loss": 0.0227, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.0284740924835205, + "learning_rate": 2e-05, + "loss": 0.0947, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.800744891166687, + "learning_rate": 2e-05, + "loss": 0.1206, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.0577640533447266, + "learning_rate": 2e-05, + "loss": 0.0982, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.681276559829712, + "learning_rate": 2e-05, + "loss": 0.149, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.6126419901847839, + "learning_rate": 2e-05, + "loss": 0.2177, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.3357105255126953, + "learning_rate": 2e-05, + "loss": 0.6325, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.11934901028871536, + "learning_rate": 2e-05, + "loss": 0.0349, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.814448356628418, + "learning_rate": 2e-05, + "loss": 0.0402, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.105903387069702, + "learning_rate": 2e-05, + "loss": 0.1089, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3973882496356964, + "learning_rate": 2e-05, + "loss": 0.0462, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.2717955112457275, + "learning_rate": 2e-05, + "loss": 0.0455, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.4529426097869873, + "learning_rate": 2e-05, + "loss": 0.1366, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.3298843801021576, + "learning_rate": 2e-05, + "loss": 0.3098, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.4207911491394043, + "learning_rate": 2e-05, + "loss": 0.2218, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.17792996764183044, + "learning_rate": 2e-05, + "loss": 0.1617, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.09431321918964386, + "learning_rate": 2e-05, + "loss": 1.0686, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.592715263366699, + "learning_rate": 2e-05, + "loss": 0.1842, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.9397964477539062, + "learning_rate": 2e-05, + "loss": 0.1263, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.34182193875312805, + "learning_rate": 2e-05, + "loss": 0.022, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.2210776805877686, + "learning_rate": 2e-05, + "loss": 0.1109, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.2350858449935913, + "learning_rate": 2e-05, + "loss": 0.0147, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.9012623429298401, + "learning_rate": 2e-05, + "loss": 0.0484, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.40518125891685486, + "learning_rate": 2e-05, + "loss": 0.0426, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.13201904296875, + "learning_rate": 2e-05, + "loss": 0.5595, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.6877580285072327, + "learning_rate": 2e-05, + "loss": 0.068, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.45389267802238464, + "learning_rate": 2e-05, + "loss": 0.0399, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.794522523880005, + "learning_rate": 2e-05, + "loss": 0.2549, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.6949991583824158, + "learning_rate": 2e-05, + "loss": 0.2512, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.1935599148273468, + "learning_rate": 2e-05, + "loss": 0.2318, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.16239595413208008, + "learning_rate": 2e-05, + "loss": 0.0088, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.11468148231506348, + "learning_rate": 2e-05, + "loss": 0.5975, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.08653571456670761, + "learning_rate": 2e-05, + "loss": 0.0649, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.5692281723022461, + "learning_rate": 2e-05, + "loss": 0.0909, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.10946527123451233, + "learning_rate": 2e-05, + "loss": 0.0692, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.266265332698822, + "learning_rate": 2e-05, + "loss": 0.0401, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.1330053806304932, + "learning_rate": 2e-05, + "loss": 0.0491, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.4198436737060547, + "learning_rate": 2e-05, + "loss": 0.0403, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.351032733917236, + "learning_rate": 2e-05, + "loss": 0.1857, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.5956119894981384, + "learning_rate": 2e-05, + "loss": 0.0257, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.6022987365722656, + "learning_rate": 2e-05, + "loss": 0.2204, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.023136833682656288, + "learning_rate": 2e-05, + "loss": 0.2789, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 6.136822700500488, + "learning_rate": 2e-05, + "loss": 1.2643, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289199494234112.0, + "train_loss": 0.1865181827545166, + "train_runtime": 96.7106, + "train_samples_per_second": 4.136, + "train_steps_per_second": 1.034 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289199494234112.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..873d6c8c70357e2d88846af5006db08fe5bf0b44 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e727b862a1870eb3ae91463e60f933be9413360f827f99704c327a06df8767b8 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..6c7c0727f5a515ec6dde256a380b5ef6d14525bb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd3796ba64746ba4859d258a9d9ed0933b70c05a088ac7c594bf2543ecc63498 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..0fd9885f9e575ac6ba5956f71d7e1ed5a7b730c5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b0f51017f78fc9872b8b04447cfa5fba94fce1478d09db70decf6973db1dc005 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..486fa3a22ce1f2f2323ce6d5c565de88ab8b3e3e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:07b14d8b2dfd20dff14c8055e0b59c26a286b07ed108c83c6a23bec65bdb2241 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..e1711d53d8fbaf15f94acce0915b09b18ee44758 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:18672e18f9b5e34bfc8c6b3ec8c999136a02ea842c62ae705916611b569d5875 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..4e4d4385126b54b2c1e748f4c5d166e276ff2d85 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aaad312a5a84b70f9e0d1fbf8e385d064ad2e2ad04ce711a406cb5ae87c79f13 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..a05ed89f43986dba2addfcbe85a4ae239604ce3a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:73ab9ee48ab80e08193acc9232611f9a87a96ab2389592d63c2610fe90b5c417 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..fab09431936e4095865f093d8119f20bfcf9b219 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3ea54248f6750e822b59bb9d761f3b41257f1cfe34c0a649a0e8cc261c5382f3 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..32b4e77f9da7a462a35287899231f7da3c7df13c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/6_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.061732769012451, + "learning_rate": 2e-05, + "loss": 0.4246, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.077659606933594, + "learning_rate": 2e-05, + "loss": 0.4101, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.9381515979766846, + "learning_rate": 2e-05, + "loss": 0.3972, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.1785246133804321, + "learning_rate": 2e-05, + "loss": 0.5082, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.954301118850708, + "learning_rate": 2e-05, + "loss": 0.5332, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.6988972425460815, + "learning_rate": 2e-05, + "loss": 0.0504, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.3482275009155273, + "learning_rate": 2e-05, + "loss": 0.3588, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.192502021789551, + "learning_rate": 2e-05, + "loss": 0.5769, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.1077444553375244, + "learning_rate": 2e-05, + "loss": 0.7268, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.341141939163208, + "learning_rate": 2e-05, + "loss": 0.3976, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.7759006023406982, + "learning_rate": 2e-05, + "loss": 0.3557, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.6827335953712463, + "learning_rate": 2e-05, + "loss": 0.2498, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.624098777770996, + "learning_rate": 2e-05, + "loss": 0.4146, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.3842098712921143, + "learning_rate": 2e-05, + "loss": 0.2711, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 4.891632080078125, + "learning_rate": 2e-05, + "loss": 0.811, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.069844961166382, + "learning_rate": 2e-05, + "loss": 0.2932, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.000631332397461, + "learning_rate": 2e-05, + "loss": 0.4321, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.8015848398208618, + "learning_rate": 2e-05, + "loss": 0.5193, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.4937423467636108, + "learning_rate": 2e-05, + "loss": 0.3943, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.40985906124115, + "learning_rate": 2e-05, + "loss": 0.3838, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.8779354095458984, + "learning_rate": 2e-05, + "loss": 0.844, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.974429726600647, + "learning_rate": 2e-05, + "loss": 0.2476, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.469043731689453, + "learning_rate": 2e-05, + "loss": 0.3853, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.4236159324645996, + "learning_rate": 2e-05, + "loss": 0.3304, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.8139584064483643, + "learning_rate": 2e-05, + "loss": 0.3626, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.3117800951004028, + "learning_rate": 2e-05, + "loss": 0.374, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.9026761054992676, + "learning_rate": 2e-05, + "loss": 0.3802, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 3.559483766555786, + "learning_rate": 2e-05, + "loss": 0.8213, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.463331699371338, + "learning_rate": 2e-05, + "loss": 0.5156, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.606950044631958, + "learning_rate": 2e-05, + "loss": 0.5916, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.52923846244812, + "learning_rate": 2e-05, + "loss": 0.3235, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.6252596378326416, + "learning_rate": 2e-05, + "loss": 0.4546, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.08812415599823, + "learning_rate": 2e-05, + "loss": 0.5337, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.3808462619781494, + "learning_rate": 2e-05, + "loss": 0.4354, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.812694549560547, + "learning_rate": 2e-05, + "loss": 0.6636, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.231074571609497, + "learning_rate": 2e-05, + "loss": 0.3085, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 6.050862789154053, + "learning_rate": 2e-05, + "loss": 0.7402, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.1351616382598877, + "learning_rate": 2e-05, + "loss": 0.5489, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.9476937055587769, + "learning_rate": 2e-05, + "loss": 0.2454, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.758622407913208, + "learning_rate": 2e-05, + "loss": 0.2275, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.4118633270263672, + "learning_rate": 2e-05, + "loss": 0.4527, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.6206480264663696, + "learning_rate": 2e-05, + "loss": 0.5723, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.9833219051361084, + "learning_rate": 2e-05, + "loss": 0.2809, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.3045551776885986, + "learning_rate": 2e-05, + "loss": 0.4316, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.9169013500213623, + "learning_rate": 2e-05, + "loss": 0.3564, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.6228079795837402, + "learning_rate": 2e-05, + "loss": 0.9473, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.5181477069854736, + "learning_rate": 2e-05, + "loss": 0.2932, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.4965360164642334, + "learning_rate": 2e-05, + "loss": 0.3169, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.5541958808898926, + "learning_rate": 2e-05, + "loss": 0.4946, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 5.342039108276367, + "learning_rate": 2e-05, + "loss": 1.0781, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 1.0467000203608064e+16, + "train_loss": 0.4613365173339844, + "train_runtime": 154.1615, + "train_samples_per_second": 2.595, + "train_steps_per_second": 0.649 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 1.0467000203608064e+16, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..145265613da21462ef6624cd23d245b8d8e41275 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fecc3201e0bb67f3abf9ababc88bc0eaf8332c35d4acfe21c115c8947b2d0d79 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..a424c0b7e588b9dcdf1bed8ebfdb2e16720dde5f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9b54330d19d93c973ff85396d9053943d423e1b46afb49dcac1d61f373bf0e47 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..f15e117b1cdf5721a353fa69f4adbce743915ee1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:44b1c4e24f4a9aa476c937198467feaa3bc3eeb01d76a1724aa8c2e47b72793e +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..6a13928c25057053fb0b88cf1d7a25f5244117df --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4cec0a5755216a1c7edba24465b25b0d4c0a1f4777ae20e09990eaa932051170 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..1718f2597aef0ad49a3cb1f0c7a328c85abed865 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2601d0d3a2104a4b6c243489d1b784e78a69b7f558c10c0d7f4cff7b61fec404 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..9813d982b2d15c61517b7de02e9d57927f45f373 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7fbe3b65c9c95b1ea043135460c261efc2451e151e36630aac1eb3f838ce83ea +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..00c5a155780cd982c565560a9631eae78c8d70b0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3828fc6e1568ca74937bb8fe45fee8ba3c26da7f8452ea0aa823aaa259055f7f +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..b9dc39f3c4ca1963a48b683a5af41f70c6f134f6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:75376701f8f83f271ffb466cc7c4d3ee68b14feb61f02a2fc72bb974cdea2895 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..5a94878fb2a6b935218121b8d8618213100b7e22 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/7_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.180075168609619, + "learning_rate": 2e-05, + "loss": 0.1053, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 5.125998497009277, + "learning_rate": 2e-05, + "loss": 0.4367, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.498584508895874, + "learning_rate": 2e-05, + "loss": 0.171, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.7644106149673462, + "learning_rate": 2e-05, + "loss": 0.2587, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.3822855949401855, + "learning_rate": 2e-05, + "loss": 0.7224, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.10668128728866577, + "learning_rate": 2e-05, + "loss": 0.0324, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.2839287519454956, + "learning_rate": 2e-05, + "loss": 0.0935, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.7780852317810059, + "learning_rate": 2e-05, + "loss": 0.4427, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.343409299850464, + "learning_rate": 2e-05, + "loss": 0.1823, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.159616947174072, + "learning_rate": 2e-05, + "loss": 0.3227, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.10373887419700623, + "learning_rate": 2e-05, + "loss": 0.2376, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.8270840644836426, + "learning_rate": 2e-05, + "loss": 0.2782, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.7245912551879883, + "learning_rate": 2e-05, + "loss": 0.0365, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.3826701641082764, + "learning_rate": 2e-05, + "loss": 0.3159, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.3875192105770111, + "learning_rate": 2e-05, + "loss": 0.3376, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.0697264671325684, + "learning_rate": 2e-05, + "loss": 0.2969, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.7924937605857849, + "learning_rate": 2e-05, + "loss": 0.0385, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.2375049591064453, + "learning_rate": 2e-05, + "loss": 0.4508, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.85701984167099, + "learning_rate": 2e-05, + "loss": 0.0667, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.222991466522217, + "learning_rate": 2e-05, + "loss": 0.6568, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.424751877784729, + "learning_rate": 2e-05, + "loss": 0.427, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.873544931411743, + "learning_rate": 2e-05, + "loss": 0.253, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.7531778812408447, + "learning_rate": 2e-05, + "loss": 0.3109, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.25735411047935486, + "learning_rate": 2e-05, + "loss": 0.0104, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.5540086627006531, + "learning_rate": 2e-05, + "loss": 0.4301, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.2908669710159302, + "learning_rate": 2e-05, + "loss": 0.4928, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.030150512233376503, + "learning_rate": 2e-05, + "loss": 0.34, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5015631914138794, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.16772501170635223, + "learning_rate": 2e-05, + "loss": 0.121, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.528695583343506, + "learning_rate": 2e-05, + "loss": 1.2831, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.6663013100624084, + "learning_rate": 2e-05, + "loss": 0.0554, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.7842200994491577, + "learning_rate": 2e-05, + "loss": 0.4399, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.6461775302886963, + "learning_rate": 2e-05, + "loss": 0.1469, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7437698245048523, + "learning_rate": 2e-05, + "loss": 0.1185, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.5476677417755127, + "learning_rate": 2e-05, + "loss": 0.3369, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.7717849016189575, + "learning_rate": 2e-05, + "loss": 0.2965, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.5303056240081787, + "learning_rate": 2e-05, + "loss": 0.449, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.4669867753982544, + "learning_rate": 2e-05, + "loss": 0.3827, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.712327003479004, + "learning_rate": 2e-05, + "loss": 0.4069, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.9958863258361816, + "learning_rate": 2e-05, + "loss": 0.8052, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.5135911703109741, + "learning_rate": 2e-05, + "loss": 0.1557, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.609668493270874, + "learning_rate": 2e-05, + "loss": 0.1204, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.0853397846221924, + "learning_rate": 2e-05, + "loss": 0.429, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.0088183879852295, + "learning_rate": 2e-05, + "loss": 0.385, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.8181565999984741, + "learning_rate": 2e-05, + "loss": 0.2777, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.7686423659324646, + "learning_rate": 2e-05, + "loss": 0.1401, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.041637659072876, + "learning_rate": 2e-05, + "loss": 0.1124, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.24596595764160156, + "learning_rate": 2e-05, + "loss": 0.067, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.4412343502044678, + "learning_rate": 2e-05, + "loss": 0.1851, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.756331205368042, + "learning_rate": 2e-05, + "loss": 0.2449, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5500391349288960.0, + "train_loss": 0.2947618865966797, + "train_runtime": 108.9589, + "train_samples_per_second": 3.671, + "train_steps_per_second": 0.918 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5500391349288960.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..c0c244800073217efdcb803cd02bdca6acee429a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:84b17a670e868d5de3aef35fba107f522c85d22ec641e893925d10cf2146b8ae +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..12315cf70ebaa09b5cac936e232ad1aa85a74703 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fa3291f071c88f6b8357150d4e55ae7b95decc087142e82ba6e2fdfe70cdba21 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..d5c77e38afd131daeaabd69a32de94465a33314b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bee775208f0c77b7ebb45a530ba25bcac50865ef8ebb3caf636b337f865ef931 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..cce43f043c0a3b00793e40211dd3d92398cc85ca --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7ab62a7e8365279a8d9bc87768346c31d9a90acf66636ae2431e2a9953619318 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..8cf362fb2dde3f532e8a3d750b4532e3897649fa --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:21b3019e9a396dfa608f119fe107b30d2535eb1f48104950478f3de1b00b893d +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..62d012aa31ef9e3a80602676e73534e65304646a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf51896b3bc6766076d7678af58941dacfa85498b7642b5ba63a696eabd75a60 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..8955498adb3377a7b3b5683f4b2909a5efe8ee1f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:52710cb82b63bebed2aafd5a0bce66114b9db0985158e58ab9356593d53ec06c +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8022e20fbf3326bab020e47f2fb1c3a50e8340c8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:53882b098dd3ea3c2804e3ff807553d1e4acdf88797f91e0f7b588169b832c3b +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..eadf435e21af87b035b55fe644406c8ea8adcf13 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/8_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.044707007706165314, + "learning_rate": 2e-05, + "loss": 0.0212, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.06966125965118408, + "learning_rate": 2e-05, + "loss": 0.0734, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.19436310231685638, + "learning_rate": 2e-05, + "loss": 0.0505, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.0242862701416016, + "learning_rate": 2e-05, + "loss": 0.2228, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.047962918877601624, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.6845811605453491, + "learning_rate": 2e-05, + "loss": 0.0778, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.047370024025440216, + "learning_rate": 2e-05, + "loss": 0.208, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 4.542006015777588, + "learning_rate": 2e-05, + "loss": 0.1953, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.390823483467102, + "learning_rate": 2e-05, + "loss": 0.12, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.210482120513916, + "learning_rate": 2e-05, + "loss": 2.2465, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.452509880065918, + "learning_rate": 2e-05, + "loss": 0.5342, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.6948721408843994, + "learning_rate": 2e-05, + "loss": 0.3421, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.5814058780670166, + "learning_rate": 2e-05, + "loss": 0.0835, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.7141515016555786, + "learning_rate": 2e-05, + "loss": 0.0384, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.2708152532577515, + "learning_rate": 2e-05, + "loss": 0.0422, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.07103218883275986, + "learning_rate": 2e-05, + "loss": 0.0085, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.359380006790161, + "learning_rate": 2e-05, + "loss": 0.7529, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.2460341453552246, + "learning_rate": 2e-05, + "loss": 0.2838, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.657318115234375, + "learning_rate": 2e-05, + "loss": 0.5956, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.625128984451294, + "learning_rate": 2e-05, + "loss": 0.6436, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 6.091564655303955, + "learning_rate": 2e-05, + "loss": 0.5434, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.0019273256184533238, + "learning_rate": 2e-05, + "loss": 0.0569, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.3364341259002686, + "learning_rate": 2e-05, + "loss": 0.3987, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.548255443572998, + "learning_rate": 2e-05, + "loss": 0.3464, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.33697396516799927, + "learning_rate": 2e-05, + "loss": 0.2214, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.04473995789885521, + "learning_rate": 2e-05, + "loss": 0.2315, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.17613528668880463, + "learning_rate": 2e-05, + "loss": 0.2886, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.6041862368583679, + "learning_rate": 2e-05, + "loss": 0.3179, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.1540800333023071, + "learning_rate": 2e-05, + "loss": 0.0665, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.08660972118377686, + "learning_rate": 2e-05, + "loss": 0.0102, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.694838047027588, + "learning_rate": 2e-05, + "loss": 0.1776, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.386035919189453, + "learning_rate": 2e-05, + "loss": 0.5919, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.136115074157715, + "learning_rate": 2e-05, + "loss": 0.1038, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.8899788856506348, + "learning_rate": 2e-05, + "loss": 0.0303, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.5832474827766418, + "learning_rate": 2e-05, + "loss": 0.0263, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.9511966705322266, + "learning_rate": 2e-05, + "loss": 0.1754, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.7049979567527771, + "learning_rate": 2e-05, + "loss": 0.0261, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.6443910598754883, + "learning_rate": 2e-05, + "loss": 0.0429, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.24294960498809814, + "learning_rate": 2e-05, + "loss": 0.1557, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.24080488085746765, + "learning_rate": 2e-05, + "loss": 0.0238, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 4.719242095947266, + "learning_rate": 2e-05, + "loss": 0.4722, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.4483841061592102, + "learning_rate": 2e-05, + "loss": 0.016, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.13410095870494843, + "learning_rate": 2e-05, + "loss": 0.0074, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.8848795294761658, + "learning_rate": 2e-05, + "loss": 0.044, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.2275062799453735, + "learning_rate": 2e-05, + "loss": 0.1288, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.576485633850098, + "learning_rate": 2e-05, + "loss": 0.4951, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.2201329469680786, + "learning_rate": 2e-05, + "loss": 0.0071, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.0893537998199463, + "learning_rate": 2e-05, + "loss": 0.118, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.22807922959327698, + "learning_rate": 2e-05, + "loss": 0.0086, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.1156950518488884, + "learning_rate": 2e-05, + "loss": 0.0076, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5327278363901952.0, + "train_loss": 0.23364763021469115, + "train_runtime": 92.4273, + "train_samples_per_second": 4.328, + "train_steps_per_second": 1.082 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5327278363901952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..ddaad89f416ecb79c57178de5ad5902a84381fa2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97357f00cd9bf8dd546d61ebbd604750e1a97d30807940e1b85228e535b187c2 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..3604e0013f9a5b322974bd521981074744103b7f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:416a7805bdf43651a7dfde3f8f0cc1a20c9d34ee6bd772ef0c3ea34bc0fe7909 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..0bb55df89825ffd3b80ab75e7ae9f8f689d35179 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:04226d0d3ef8cc3acb31098d21bdf5a8926453f42f10f8c7810eb0ffbdfb2f56 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..7f1c70cda1c23be38266b2a816e9304cce2efeda --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6a95fcbda7b5e2e1df0ee2acc233d741b6b9a349aae8ec2d789179034bcfc9cc +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..880d17c9b66f77f7599be36531bd554a89326a9c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2def69b2591bc167ab4a390e37198442e2477df6d1279bc7fc3ac2c512a9c42f +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..60847f08647f56813d59eb8b25ffcd930bbf5c0a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:da7e49ccdc081e87e33b5354320591f86f1863e700b6e4880d8a9900673724fa +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..4e397db44a38a95ac171a23b21c177770a21ca80 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2d455c34b8109dbb820af84a7e100e0ef376a2f59bdf31bda047dcb566bb0ab0 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d437340a815520258366efb89510a96df1a1647b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:861e2164d56659cc9e9985c895ce3adc135849ee1869ab303fd78b3319f14c94 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..55ff931362dde4c8e19219b2e8d051ec0d45d47c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/9_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.3204832673072815, + "learning_rate": 2e-05, + "loss": 0.3541, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.7209656238555908, + "learning_rate": 2e-05, + "loss": 0.6598, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.0757590532302856, + "learning_rate": 2e-05, + "loss": 0.0849, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.2861272096633911, + "learning_rate": 2e-05, + "loss": 0.1144, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.22115656733512878, + "learning_rate": 2e-05, + "loss": 0.3237, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.745354175567627, + "learning_rate": 2e-05, + "loss": 0.1459, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.0825243592262268, + "learning_rate": 2e-05, + "loss": 0.3776, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.434133291244507, + "learning_rate": 2e-05, + "loss": 0.5632, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.9266154766082764, + "learning_rate": 2e-05, + "loss": 0.148, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.162006139755249, + "learning_rate": 2e-05, + "loss": 0.238, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.5716254711151123, + "learning_rate": 2e-05, + "loss": 0.6949, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.822958469390869, + "learning_rate": 2e-05, + "loss": 0.4768, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.7830159664154053, + "learning_rate": 2e-05, + "loss": 0.4661, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.1483240127563477, + "learning_rate": 2e-05, + "loss": 0.0667, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.2384364902973175, + "learning_rate": 2e-05, + "loss": 0.2142, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.132415771484375, + "learning_rate": 2e-05, + "loss": 0.0051, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.9311197996139526, + "learning_rate": 2e-05, + "loss": 0.5334, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.0047414302825928, + "learning_rate": 2e-05, + "loss": 0.1647, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.26010996103286743, + "learning_rate": 2e-05, + "loss": 0.0543, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.6168017387390137, + "learning_rate": 2e-05, + "loss": 0.3376, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.2097108364105225, + "learning_rate": 2e-05, + "loss": 0.3008, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.657468318939209, + "learning_rate": 2e-05, + "loss": 0.592, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.2511478662490845, + "learning_rate": 2e-05, + "loss": 0.0745, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 4.062618732452393, + "learning_rate": 2e-05, + "loss": 0.4644, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.859788179397583, + "learning_rate": 2e-05, + "loss": 0.0759, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.9490105509757996, + "learning_rate": 2e-05, + "loss": 0.1003, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.7533918023109436, + "learning_rate": 2e-05, + "loss": 0.4229, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.20891602337360382, + "learning_rate": 2e-05, + "loss": 0.0147, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.801180124282837, + "learning_rate": 2e-05, + "loss": 0.3486, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.9825308322906494, + "learning_rate": 2e-05, + "loss": 0.1038, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.1416311264038086, + "learning_rate": 2e-05, + "loss": 0.187, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.481233835220337, + "learning_rate": 2e-05, + "loss": 0.2438, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 5.533125877380371, + "learning_rate": 2e-05, + "loss": 0.2529, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.0692771673202515, + "learning_rate": 2e-05, + "loss": 0.091, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.15696126222610474, + "learning_rate": 2e-05, + "loss": 0.015, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 5.583909034729004, + "learning_rate": 2e-05, + "loss": 0.4656, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.8694599866867065, + "learning_rate": 2e-05, + "loss": 0.2253, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.963747262954712, + "learning_rate": 2e-05, + "loss": 0.3493, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.33902794122695923, + "learning_rate": 2e-05, + "loss": 0.1497, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 4.48232889175415, + "learning_rate": 2e-05, + "loss": 0.4243, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.707913875579834, + "learning_rate": 2e-05, + "loss": 0.2604, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.57442045211792, + "learning_rate": 2e-05, + "loss": 0.1924, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 3.4521143436431885, + "learning_rate": 2e-05, + "loss": 0.2211, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.3788692951202393, + "learning_rate": 2e-05, + "loss": 0.4068, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.4747298955917358, + "learning_rate": 2e-05, + "loss": 0.2605, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.7445545196533203, + "learning_rate": 2e-05, + "loss": 0.3041, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.118432521820068, + "learning_rate": 2e-05, + "loss": 0.595, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.7435587644577026, + "learning_rate": 2e-05, + "loss": 0.0552, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.2370156049728394, + "learning_rate": 2e-05, + "loss": 0.2457, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.26646748185157776, + "learning_rate": 2e-05, + "loss": 0.0594, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291656215527424.0, + "train_loss": 0.2705048847198486, + "train_runtime": 95.3362, + "train_samples_per_second": 4.196, + "train_steps_per_second": 1.049 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291656215527424.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7af43b43b1f3b391ca773005deb77d6eaf2df9eb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:37621c4324b90d7fd65bf9c591951cb2bcefbef1fcecda13644596b1420dba7b +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ca6ecf93b6caedf580079c283dae42c9dddf3154 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2802479273d49fb9acfdf24f472f86142a081a7d208a424e8d018796322f4d8b +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..19ed20bd6362c9e969d674eb89cab935458714c7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d7e335d71fe4344bd3c0f4bddbc7d7d4287b0a5db988febf9acb49c75d8e46c +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1650ff7aee694395f543de5cf454a081c5eb3920 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7dc7b7e3509fbe55497adf78df2f11036f80a1f5cfb84294e4be1bac27d48edf +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..04785518d86cf82a2cbe2e0e3e9a5d7982d847b5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9ab865a6fde8eaccbbc7684bc85f94667ca023b8bc391cc6ed48392bfa3cee42 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..74c588f1a6f9902ec68c67a8d72f4b852fbe9021 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b0f13d3a6d0abbb747e70b94a11f427df24e69ddbc2c3e1d289ce11b74a90f9f +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a16d626677ddcd10c5ee94f3c34e9b5952a746f6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b9a953bfda8f60e1dc85231a8a336259358e856aa9026906eb064b06c2cb1dd +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b3af175be5f42014abbfc5609e5e0fdd721ffa8d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c45610a2e05df8c4dc1e02db3729ba5e5003949d6cda284d95609aa5ad61cabe +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e6a4489f90ec71371cbeec2a74413a1cff93b45a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:42d4af201b8ccf51dc5f36ebaf34c89e8673ef3d11850bf840f33b5b672a3b63 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e61d87f0b98230b38fefc123653ec5afe54713d0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d583fc65f2a490dec35b64279f2020284eba59ed27011883d88c5acaf013a7d8 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4423576bd4b9b852caccbbf9c31f266960209d13 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8e7c1e68ff1565b6a0faea4a1913a583700db881eaab48dee2c0e20bc47ce2f8 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a51473dc225f3996f7d95e7f9f83fb0b150b27ab --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1a9a4d58ab02d9df542469da13a3e132234a1f7240618dc45486b52f3518254e +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f1a307b59359318fa39bfdffa43b04b4c4c07ef2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5350b67d4c3b00b6d59dc0aecda28cb0372cfb82e027027e55ba22aec1d1bf48 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d78f41b10aa6063a8089fd36b63a4969e3e1da6a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:83d63573884d54134a4e75eca7d7fe6d2d10c8a4d8ac4b2c60d529ab954ed1fd +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..55fdbef98ed5fddcb43833c0bbf3d4ba84741da4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:91490259fb922b65c8d5efb8e9eb9036393c913a25e4f51ff7fb32d24f017aab +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9084da734d8aaf928a83946f1ecebf84349fe141 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b69aafa0f677ae75b0344a39f5427c59a2f68cec076f2b9ee980765c8a4bd5cd +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..856a7a3af614c8a3c5a20940e3d5b36b3726d24c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:96576983026c49bc36f1f39ef342d4338393555982aa7f35083eca0e199e0851 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..13466e91f82de567ef4a5e2477ce7ec6d689fe3c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e3e5d028589c4dbacbdf0161d6e6fd605162fd4faabaee1179f1aff147fbff7 +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..298d2820f233dfbb5aa3f0051ef9fed6208dd3ae --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4daf66a62a129811b5610a22ca25dad2b8ba8a6fde2f55239e0aec6916473247 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..283c33eb01892338635852fc9b17c577df64a8b6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:506020a3ff4e9cfa87847df0ad61eb59e9a3ad9adb6827871a92540806860f42 +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..8a5da84ba8d261a62f30e86aff281272970eab42 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0d5b608994ec07d706f2c95253fa2f34ead58a2b7b01743073244987a377caee +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4e1b915b62b1752537f2ceb67622852f4b2d73e3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8053d5663b90a01d9a109d0375c7d994725c723ab39d1648f8ffd6350ae7689f +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..cd2739a1ea6a0ff3fa8fa95e4d1531d17c7a0052 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b26e38306b66ae76d2e0c6140b2bf9537c39dee8b981f0597338766fb8080cc +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c92c62b5002fcf34f43e4690cc666467db62eb6c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:af32ce326c2d9becf4c580595fcc93563d9aa6cf4147efd400dc0dfc5fc528ec +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..582ba05819266252b014837ede4f3ad1122904e4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:221376374e96aa79b9e95efeb78671c0de63b9c86b575863115ba77fd8621d7a +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..3feb851c2267382a0f6c1263ca82d8cda9400a3c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c6ab0fb9b6477fb042d65d96f98c315a6d2eb4772494c080cb02b2a75a4d641 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..bc8efdcc2d6406d520f1dcf6f14ecfa5c15e308b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f4ebcf76e4cb423048282d1f68a26b73841cb889ea407add949065a00f3d35af +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..bec22f50f3abb569e01d186a786510e59b078d94 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:24b496fa6aec1b492f1c4f39669b3df7e435acb0dfd17b12a58853d229e239cf +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0c10e9ce21e69e9d1fd09e25cdc91dc42b789046 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a14b9c83555350b33600049e5b4be87d33ba0393613b17697630bbc113daebda +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..bc3ee88c671e471273073a26b343cf688afb089f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:762aa18754399082e3f3160052613bf6ffd86538b7debf0b3c1eee013ffb734b +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1e554f3266a28f86cf83923d830fb033d8c9383f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3bb13f645ee377cb89c32f155c452952bde668650e0156c999c2f4a7c1393342 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..88acacb539361400a4d92ac4e9348e9f36a9ca38 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d896541be0f966b08a0eef6e0a50b85b892fe760037093166a7e593a5d6ea65f +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..37401cd24941cb268ae99153af4ee3265d6bfc66 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:239e437a0b6b74a3e21202938a21d9b410035153134538a9330fb0e3b5dc8a31 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..64b92dbfd1631dacc473d97475cf282d5fcf0074 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ec5b37e3cf73428e7b7dc614d423cd502613f0b61e5d354c6d18982db641a98d +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..eecabd53687883b88a666b43b4df224cb19befb5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:86a2018eefd337066b53aeab9f97af2a149464d4bf614fa8d17f7e7f9bf5efad +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f02cb946925c4a05c5098e51eca01eda09aa090c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a53eff4e815195fc05014d0bccf4f605f10129bcf19a61afa3a2735446a23b0e +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..da9ddbade8fa2c453275a40896ebb2ffb18f4d46 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:69a74214355071ef8b97444a7a315a5919755c7d45859ab0430a4ad3ea07fba6 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b057058c25f5fddc7a305e1f9fa8dbc97fc2114d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:58a507af3f0171b784414063cab46e1095d23e6f8f72362383f7f7497fe3ac7d +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..36e001018759215f65d427c387c4b8f5004be4ef --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc4e0710303302e4c9a4544ea35f608f58a2f860d88497b056cd1e84b6fc5b4d +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..822ca90207aaa4165076f9c4d5db216767e85692 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c2f9c4f0f443c6c39c26e39baca7faa71a0df230172479a3bb758b7d7b082e7b +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..8fdeadb8e47eba19a1d55b90df5d3b76c58690a7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9c5eeaec4b4cdb890934ebdfb44912d1b4eda71d82d5c65b3457642a0789abe8 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6027a48309b25e5274aa46a2be61f21a8ba4996c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:468dadcb7d8c037c8d25bb5a68c852f77d94ae34d64d48b65b69859ba6ffe953 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round10_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round10_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..8480f3290fe02dcc8ad3f411dce4bb771607f136 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round10_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2a7a8d47edc80b205a0230dd0d312c16f3cc81aa0b01bd0d21718b6314735bb +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round11_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round11_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..120369764554a4d6cac58078d8512bf359607829 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round11_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf1dde35905b1ad125edf778aff1f9c3576ae0715c8c9876f665a63efcf18cd +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round12_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round12_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..64b454d945aa2a842c509c5cebe5ce0706d83a0c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round12_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb23ed1efc2080025ddd7ca376ff771b3f058afa24e048f5cb3cf99201aa1e03 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round13_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round13_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ad8d07954c41361e73a0c9c7d3b4b46d894b188 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round13_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8caca9dd439d57dbc5054f65d33ebbd7ac694940b2b9dbe5e8cdf68881137a2 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round14_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round14_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f05da1a0577050cea961f470b13842ab5f1d3f9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round14_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafca5706ac84f575d7648ee4786543561523e58cea8c01d2214b03d300bb988 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round15_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round15_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..351a4d2a284717aec90435f50067e1bb7002e9e4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round15_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:341a5a225c4c17a621d4f60caf4c34834000ea74e6f3c8b4f5aaada72ca1f547 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round16_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round16_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..08cfb19e58b68e297c605b87165d4acd241054e5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round16_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb378347913e2dd7be50fa6f9e3bd9e481f3897c2c9d31c597fe9164efc50bee +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round17_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round17_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..aeee08397c6123d33648f6703464379a4b7f85fe --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round17_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70c677ae195453774b8815440fee856ddb773c322681ac422ebc6e8abd5a58d7 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round18_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round18_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9935621fdc6da9e4abad379d97303b5ebf4f8f5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round18_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:393c3d5a15b40b2ff049c254f5b09f25d226b0971827b4e7460ca58f9ff7b76f +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round19_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round19_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..2766cf1a4da1390df01ea4146b4a3a7a06e1434e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round19_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dc77db94ac5b4746eebf2c850a9425da9a425a71108790d6b7f5f2654b614a8 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round1_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round1_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..34bfc5121e29fc587c3f574c17d8006189d3021a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round1_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47a94b88b812f21101345aae20a53cc10e0b346660c2cef669c77b201d06d508 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round20_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round20_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..4be93ba48f85d89c8cbd0e7426b79a10f51bc786 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round20_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20b9a2e67fdcb4051cebf52cba2ed0d47facb4dc69170f35a3483e3601bf25a4 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round2_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round2_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..b42d717305b30cd186c1c408b8719020c0658a5d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round2_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fd7458877767da0a36b80344cd6ec7eae8fbc8e11d07f8bfce3b4d5f5f239a9 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round3_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round3_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..01842d1316f21c3326c5735c6257b499da7ca458 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round3_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e537f056156082569ae0819abd06beb04b46581d0c63843bd68926e0ac2bff0e +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round4_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round4_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..3af9fe115499486166c3ef6d0a7bd27661ac08b4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round4_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a9649d6997f84508e5f566393e91e42e03a31aca55878bacd7ffd2e134bbfc +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round5_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round5_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..abd1e557c23a8769870b0beedaf7df93836a101e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round5_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7520bace4f501771d879621facd1b44f8db1ad7385a653e1535c7642465e67d +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round6_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round6_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..ebc846b7a5c587ced387cca85c4065e921c8cbe3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round6_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1951a910d1b7587459db2110e8ac43cbc906a34267372eb21d2244a900eaec46 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round7_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round7_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..5eb1370d3fa3177cb4fb710882f4b7eb706817e8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round7_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03f40a0df2437812b9e6de4e5d4e30014acb71344efa8132f19e1fc0f5925c82 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round8_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round8_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..385c599585a0973f240a027d28e985ca72d6b721 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round8_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3ef6d8247bc37b017895c9e5bbaf32f2db500f94f66b4ebade5a4a5ba7a271 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round9_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round9_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..219adb8927937a4597cd82efbf7ebb4c8a7519ad --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAggOnly_iter100_round5_hetero/round9_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:277fa5f0607d993e80153171e4760f999c4ee6a65452ea62694d0b9b03818f22 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..2b993a2d80bc51404b56de3d2360c09629c68117 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2c59b9a0c42d61aa6138dadb39431adaa9e3f41574f27b0641e63e9716883541 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..94adb839ba8a293aaa64131b634bef5cb8aa7f18 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9397599509cf91284b1abfa4a30f69664b3f95b90e93a1487d60e1ca085399e4 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..c3ff110b3dc1d38b02ae44f462143b9a6ba20a8b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:39ed26f142b50e1e3a4da0614532b9006bdbe713ac7fe327c90859131876efa9 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..8c04d1146635e9bc149fa7ccc30c49da4558ccba --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e1e6bdbdcff4b9b02275f617e2df6466c4d384ff970307e38a9d34076087c277 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..a41086c440c170adbaff54d58df88145be3ec51c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:020b4488a35e5781506ce7eeac6b12036ed5304e91d7f67bc1766a5493222781 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..7d83f18cd119fb8a5da58977270fb282949de07b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d51fe8f33bf89513fa84d8acad560cc5b98d38f110777e520ecf67f72f4da869 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..ea98fccf06699ca7d37c4f6a22c0a10e5e169ace --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e9bf2afe513f03ddf951b620a0179ac5abfa3f913142cb1b026ccba374460bb +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..0e63d8ae6aab16a2cac0d49218f4377caa84ec5d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d0f36349e58b053268d9da93ddb1f23be232a496a735973325e950a979aabe9c +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c4d235d1460d7050a92e405e7acb2c36cb96baed --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/0_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.0181021690368652, + "learning_rate": 2e-05, + "loss": 0.2124, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.1943154335021973, + "learning_rate": 2e-05, + "loss": 0.1898, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.5280537605285645, + "learning_rate": 2e-05, + "loss": 0.036, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.24016790091991425, + "learning_rate": 2e-05, + "loss": 0.0396, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.218228340148926, + "learning_rate": 2e-05, + "loss": 0.2633, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.7464208602905273, + "learning_rate": 2e-05, + "loss": 0.2551, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.2280319929122925, + "learning_rate": 2e-05, + "loss": 0.1021, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.9998395442962646, + "learning_rate": 2e-05, + "loss": 0.4661, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 5.5077409744262695, + "learning_rate": 2e-05, + "loss": 0.9844, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.448251962661743, + "learning_rate": 2e-05, + "loss": 0.3192, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.18104611337184906, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.8170197010040283, + "learning_rate": 2e-05, + "loss": 0.3799, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.22930094599723816, + "learning_rate": 2e-05, + "loss": 0.1744, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.29661527276039124, + "learning_rate": 2e-05, + "loss": 0.208, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.7264126539230347, + "learning_rate": 2e-05, + "loss": 0.2829, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.11997828632593155, + "learning_rate": 2e-05, + "loss": 0.0252, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.077545166015625, + "learning_rate": 2e-05, + "loss": 0.1558, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.941800117492676, + "learning_rate": 2e-05, + "loss": 0.3011, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.714742124080658, + "learning_rate": 2e-05, + "loss": 0.229, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.3819150924682617, + "learning_rate": 2e-05, + "loss": 0.1009, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.2275683879852295, + "learning_rate": 2e-05, + "loss": 0.0173, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.9439992904663086, + "learning_rate": 2e-05, + "loss": 0.2885, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.7579358816146851, + "learning_rate": 2e-05, + "loss": 0.1088, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.32029101252555847, + "learning_rate": 2e-05, + "loss": 0.2027, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.1474190950393677, + "learning_rate": 2e-05, + "loss": 0.0823, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.2990169525146484, + "learning_rate": 2e-05, + "loss": 0.0952, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.36493608355522156, + "learning_rate": 2e-05, + "loss": 0.1531, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.109502077102661, + "learning_rate": 2e-05, + "loss": 0.1239, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 4.559152126312256, + "learning_rate": 2e-05, + "loss": 0.5876, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.10711327195167542, + "learning_rate": 2e-05, + "loss": 0.0061, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.5605273842811584, + "learning_rate": 2e-05, + "loss": 0.098, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.5725798010826111, + "learning_rate": 2e-05, + "loss": 0.1081, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.5068108439445496, + "learning_rate": 2e-05, + "loss": 0.0177, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.3335644006729126, + "learning_rate": 2e-05, + "loss": 0.1052, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.9281625747680664, + "learning_rate": 2e-05, + "loss": 0.2438, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.601179599761963, + "learning_rate": 2e-05, + "loss": 0.0844, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 5.78809118270874, + "learning_rate": 2e-05, + "loss": 0.9488, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.186406373977661, + "learning_rate": 2e-05, + "loss": 0.4568, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.5558359026908875, + "learning_rate": 2e-05, + "loss": 0.2959, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.6190751791000366, + "learning_rate": 2e-05, + "loss": 0.3141, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.7734168767929077, + "learning_rate": 2e-05, + "loss": 0.0784, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.545628786087036, + "learning_rate": 2e-05, + "loss": 0.1852, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.8479626774787903, + "learning_rate": 2e-05, + "loss": 0.0373, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.5434174537658691, + "learning_rate": 2e-05, + "loss": 0.1243, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.232873916625977, + "learning_rate": 2e-05, + "loss": 0.1415, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.8143606185913086, + "learning_rate": 2e-05, + "loss": 0.1871, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.973983287811279, + "learning_rate": 2e-05, + "loss": 1.1553, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.6529728174209595, + "learning_rate": 2e-05, + "loss": 0.484, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.042940378189087, + "learning_rate": 2e-05, + "loss": 0.3218, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.31299078464508057, + "learning_rate": 2e-05, + "loss": 0.0614, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289495989583872.0, + "train_loss": 0.23702473640441896, + "train_runtime": 95.2875, + "train_samples_per_second": 4.198, + "train_steps_per_second": 1.049 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289495989583872.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..209427e303ac52c437ecfe44f33349b83948f20d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f320ae4a6b4da89fd2547e59015def039fc44ebc400a0fa66e259524b38c0c7 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..7127df6cdaae38ffa70c62a341311e11cefea9ca --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:611d26ee241e094399f003a804698c88ab3c2b76d25ca8b63355aa927c879124 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..8c6ab8f87204443c536a7f3e05bdea9909d70238 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9489ea300b7aae94c663d80fc11608090c4a2facdb18e35713bc26842a591a16 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..8f11b81117cd3d572732fd106767302507b37feb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:79f96df86455f63ad6e62c9a95489464733f69ae41dcee6c4bc711b39b0b5e60 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5bbf7b6db23d685681d7e9eec1d080ef65a73c6f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fd76b95b9f2e93a47bfc9221257b3124f7de3e8e73a86bbab6461ce39257737f +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..1bbcfb80a1190056dab7016ddc4932edb0e116f7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8df34fb4ebc5a5f7895f8afea277fd768a9e2c4a4d2a5f7afee7834ce0854a20 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2b05b291d96476c8e8660eb5c180ea5913c6124b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d2b8b369e7e8a3557451c6815b32f173ba58e01cca2af64190dddb8670101100 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..bf4d41d415e94e99b259cf4becd77e4ad67c632f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a4775866b2ded9f446395aa77ce31e69352730265ce3166a68e1375383e8e297 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c81d2f484eef209d0de466d00b0391eb2a3552ee --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/10_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 5.721604824066162, + "learning_rate": 2e-05, + "loss": 0.3893, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.846872329711914, + "learning_rate": 2e-05, + "loss": 0.1764, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.260456085205078, + "learning_rate": 2e-05, + "loss": 0.1262, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 5.335000514984131, + "learning_rate": 2e-05, + "loss": 0.5851, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.3808627128601074, + "learning_rate": 2e-05, + "loss": 0.3073, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.2647400200366974, + "learning_rate": 2e-05, + "loss": 0.0152, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 7.55369758605957, + "learning_rate": 2e-05, + "loss": 0.4288, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.3966429233551025, + "learning_rate": 2e-05, + "loss": 0.1647, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 4.041947364807129, + "learning_rate": 2e-05, + "loss": 0.1163, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 5.1376237869262695, + "learning_rate": 2e-05, + "loss": 0.305, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.986076831817627, + "learning_rate": 2e-05, + "loss": 0.2444, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.254326820373535, + "learning_rate": 2e-05, + "loss": 0.1503, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.809089660644531, + "learning_rate": 2e-05, + "loss": 0.5638, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 9.87086009979248, + "learning_rate": 2e-05, + "loss": 0.6096, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 9.749313354492188, + "learning_rate": 2e-05, + "loss": 0.5536, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.0193254947662354, + "learning_rate": 2e-05, + "loss": 0.058, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.975440740585327, + "learning_rate": 2e-05, + "loss": 0.1336, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.24320946633815765, + "learning_rate": 2e-05, + "loss": 0.289, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.644618988037109, + "learning_rate": 2e-05, + "loss": 0.1635, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.8331212997436523, + "learning_rate": 2e-05, + "loss": 0.3298, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.46929749846458435, + "learning_rate": 2e-05, + "loss": 0.0699, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.23767855763435364, + "learning_rate": 2e-05, + "loss": 0.6411, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.2631480395793915, + "learning_rate": 2e-05, + "loss": 0.0512, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.7858497500419617, + "learning_rate": 2e-05, + "loss": 0.2562, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.255497455596924, + "learning_rate": 2e-05, + "loss": 0.1045, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.648743152618408, + "learning_rate": 2e-05, + "loss": 0.9495, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 11.582180976867676, + "learning_rate": 2e-05, + "loss": 1.0654, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.41179004311561584, + "learning_rate": 2e-05, + "loss": 0.0161, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 7.332399368286133, + "learning_rate": 2e-05, + "loss": 0.1339, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 8.741049766540527, + "learning_rate": 2e-05, + "loss": 0.3881, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 6.80659818649292, + "learning_rate": 2e-05, + "loss": 0.4437, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.2808878421783447, + "learning_rate": 2e-05, + "loss": 0.1128, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.164119243621826, + "learning_rate": 2e-05, + "loss": 0.7794, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.4106810092926025, + "learning_rate": 2e-05, + "loss": 0.3718, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.4070886373519897, + "learning_rate": 2e-05, + "loss": 0.0445, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 14.359206199645996, + "learning_rate": 2e-05, + "loss": 0.8547, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 5.361375331878662, + "learning_rate": 2e-05, + "loss": 0.2578, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 9.94933795928955, + "learning_rate": 2e-05, + "loss": 0.3855, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 6.499366760253906, + "learning_rate": 2e-05, + "loss": 0.7686, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.415165662765503, + "learning_rate": 2e-05, + "loss": 0.6794, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 7.7353057861328125, + "learning_rate": 2e-05, + "loss": 0.3724, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.08788453787565231, + "learning_rate": 2e-05, + "loss": 0.0083, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 6.590057849884033, + "learning_rate": 2e-05, + "loss": 0.3871, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.4260687828063965, + "learning_rate": 2e-05, + "loss": 0.5413, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.1965910196304321, + "learning_rate": 2e-05, + "loss": 0.1441, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.839502811431885, + "learning_rate": 2e-05, + "loss": 0.4268, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.330932140350342, + "learning_rate": 2e-05, + "loss": 0.6973, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.8875535726547241, + "learning_rate": 2e-05, + "loss": 0.1165, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.881289005279541, + "learning_rate": 2e-05, + "loss": 0.2626, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.7781264781951904, + "learning_rate": 2e-05, + "loss": 0.4982, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2220576018006016.0, + "train_loss": 0.3507712364196777, + "train_runtime": 74.7806, + "train_samples_per_second": 5.349, + "train_steps_per_second": 1.337 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2220576018006016.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..8c3373c68ab80b3e31757e71f757b92ce876c484 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6000e779f867363de877e77d325060684b42d07c967dc1af09d67e010c45a94 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..c941c551a8ae9dc8553c028db00cb31e5ffd19b0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0601583c39ca0da0fe650b4587c377f161302f6d2f427cbae0c7d27db9436b66 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..ff81073ef268eba2d7a354aa2a86ee78f7cb726d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7dd298d0888c436f64da49f6f16a1262cd0c88f1f68a1497afe61bd6091e1cd +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..66021570441d1a772ffd5322177aa51f42b8f7b6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b96fccd3a43b31b551139ad5d64cc969c8f90384a7e9c14f2af929a321ddc338 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..746b100e0f86eeb09fce50ffebe6cd7c4de9b30c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f5def677bc79c7986ec01d21125de844c5bc86d60e19b4e7ef913e99516cfb7f +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..880e05bd9670ab37ed62514b421a95c1c1308774 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:57d0d5c4f847885fd0408a67a8c2c67cc08d581487fa91d18f960b9af2a2bfd1 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..dd2b447c9fde4213816cd6b211419001457fe34f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e977da2f26137bb5c3f394b40e44e39bfffa4fbf498deb334aa2ea17e72bb346 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..bae42c366af90324355b956a3067ffb65447961c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:154839faa81e177a38e9c2cbfef1b79ee3d17383698161ca4d55fce93b3342d4 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..090547bf9b125ebae03b9bbbf54fd09c83736b46 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/11_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 7.610400676727295, + "learning_rate": 2e-05, + "loss": 0.5386, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 7.33369779586792, + "learning_rate": 2e-05, + "loss": 0.4753, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.2170541286468506, + "learning_rate": 2e-05, + "loss": 0.3598, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.8871169090270996, + "learning_rate": 2e-05, + "loss": 0.3388, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.3350019454956055, + "learning_rate": 2e-05, + "loss": 0.6953, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.37825083732605, + "learning_rate": 2e-05, + "loss": 0.4648, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 5.120723724365234, + "learning_rate": 2e-05, + "loss": 0.9126, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 6.384683132171631, + "learning_rate": 2e-05, + "loss": 0.7183, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.0490331649780273, + "learning_rate": 2e-05, + "loss": 0.5811, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.9757301807403564, + "learning_rate": 2e-05, + "loss": 0.4363, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.0690664052963257, + "learning_rate": 2e-05, + "loss": 0.3354, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.3635153770446777, + "learning_rate": 2e-05, + "loss": 0.4695, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.3835809230804443, + "learning_rate": 2e-05, + "loss": 0.3666, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 7.509660243988037, + "learning_rate": 2e-05, + "loss": 0.4133, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.878528594970703, + "learning_rate": 2e-05, + "loss": 0.5303, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.9740195274353027, + "learning_rate": 2e-05, + "loss": 0.437, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 4.0781049728393555, + "learning_rate": 2e-05, + "loss": 0.521, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 5.338316917419434, + "learning_rate": 2e-05, + "loss": 0.5234, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.748468279838562, + "learning_rate": 2e-05, + "loss": 0.509, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.816670298576355, + "learning_rate": 2e-05, + "loss": 0.433, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 6.616994857788086, + "learning_rate": 2e-05, + "loss": 0.3572, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 5.785435676574707, + "learning_rate": 2e-05, + "loss": 0.616, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.8161792755126953, + "learning_rate": 2e-05, + "loss": 0.491, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.073464035987854, + "learning_rate": 2e-05, + "loss": 0.257, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 5.716798305511475, + "learning_rate": 2e-05, + "loss": 0.6142, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.189061164855957, + "learning_rate": 2e-05, + "loss": 0.5015, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.811913251876831, + "learning_rate": 2e-05, + "loss": 0.3814, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 4.027775287628174, + "learning_rate": 2e-05, + "loss": 0.2551, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.067948818206787, + "learning_rate": 2e-05, + "loss": 0.2354, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.7939854860305786, + "learning_rate": 2e-05, + "loss": 0.2966, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.3211474418640137, + "learning_rate": 2e-05, + "loss": 0.3953, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.8106167316436768, + "learning_rate": 2e-05, + "loss": 0.3467, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 6.369966983795166, + "learning_rate": 2e-05, + "loss": 0.5654, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.12884096801280975, + "learning_rate": 2e-05, + "loss": 0.1759, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.613821029663086, + "learning_rate": 2e-05, + "loss": 0.3548, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 3.5542542934417725, + "learning_rate": 2e-05, + "loss": 0.4522, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.0574951171875, + "learning_rate": 2e-05, + "loss": 0.6187, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.27023446559906, + "learning_rate": 2e-05, + "loss": 0.3988, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 6.523299694061279, + "learning_rate": 2e-05, + "loss": 0.7523, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 5.243956565856934, + "learning_rate": 2e-05, + "loss": 0.6826, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.265439748764038, + "learning_rate": 2e-05, + "loss": 0.8018, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 4.8929667472839355, + "learning_rate": 2e-05, + "loss": 0.7393, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.5818774700164795, + "learning_rate": 2e-05, + "loss": 0.2957, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.911024332046509, + "learning_rate": 2e-05, + "loss": 0.49, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.023634925484657288, + "learning_rate": 2e-05, + "loss": 0.529, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 5.377819538116455, + "learning_rate": 2e-05, + "loss": 0.3633, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.495153546333313, + "learning_rate": 2e-05, + "loss": 0.4297, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 4.199329853057861, + "learning_rate": 2e-05, + "loss": 0.4282, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.6200451850891113, + "learning_rate": 2e-05, + "loss": 0.3918, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 7.272134304046631, + "learning_rate": 2e-05, + "loss": 0.4976, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2191347477905408.0, + "train_loss": 0.4754704189300537, + "train_runtime": 62.5311, + "train_samples_per_second": 6.397, + "train_steps_per_second": 1.599 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2191347477905408.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..3ad6e8b7ca61f14e725688f1913d5f091236133c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cf37cae3456badd461be657af7df2c48f8d82b190670900fe51a6ffe8ca7dec4 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..9cbdb4b731c707546b8d561571481bc1730bf43f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a859c6c9dc9c5b35fecb88014dd7d549d6f390d2d0f18a8a9e9280caaa61d9ee +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..c1742ca59e992614b4b712cfea01bee848ad7b17 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bb4150860e1270f2f44bce2ebd0d6f9dde51b38bfa586f94cfd74a3249680668 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..39f8573f06921f4e83596b3e6f6a9aa634f41230 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e110c641949d978f6f65b69c6fc3260f08fafd20579349e37819031acb2e4d24 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..96f24c0bcd3163527234b9dc87734f2369b5cc0c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8bb59f9309d22ef261b1e09fa4128997381576e5c6294c5113543adc1e0fd1f +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..7c2dc90cf197c1c922350c22cb162439f84742e1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7ecd8122b9ec692fd973f6c613d55ae553ad48e30b45f332b24267d8603e8e73 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..6cd17f74de99ecd3e5cbbcc24ae4f00974b3d0a6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0c6674b6bc178941c6a572785b1663297be80709f905650b130a1e9a3ea7a55e +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..0c12e1a92eac91d15941ada31283011211e7dbd9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:188790402204e1209e1d13f9efe4eccc1ec082b1033bb80e53a48d297a4a36ee +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..1f36310368ac4b038fef1ebf5f66c082df62936c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/12_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.284913182258606, + "learning_rate": 2e-05, + "loss": 0.0404, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.0025015019346028566, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.06335554271936417, + "learning_rate": 2e-05, + "loss": 0.0152, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.7913476824760437, + "learning_rate": 2e-05, + "loss": 0.0266, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.3213243782520294, + "learning_rate": 2e-05, + "loss": 0.0169, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.09720367193222046, + "learning_rate": 2e-05, + "loss": 0.003, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 4.3638811111450195, + "learning_rate": 2e-05, + "loss": 0.6127, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.7336924076080322, + "learning_rate": 2e-05, + "loss": 0.2785, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.4557226300239563, + "learning_rate": 2e-05, + "loss": 0.144, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.775259256362915, + "learning_rate": 2e-05, + "loss": 0.2038, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.17942756414413452, + "learning_rate": 2e-05, + "loss": 0.1056, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3040379285812378, + "learning_rate": 2e-05, + "loss": 0.0159, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.42362678050994873, + "learning_rate": 2e-05, + "loss": 0.0556, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.426983594894409, + "learning_rate": 2e-05, + "loss": 0.1332, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.017152803018689156, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.06773574650287628, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.10270370543003082, + "learning_rate": 2e-05, + "loss": 0.0343, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.17765872180461884, + "learning_rate": 2e-05, + "loss": 0.0747, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.22427186369895935, + "learning_rate": 2e-05, + "loss": 0.0229, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.8372613191604614, + "learning_rate": 2e-05, + "loss": 0.0455, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 4.128047943115234, + "learning_rate": 2e-05, + "loss": 0.3322, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.029971778392791748, + "learning_rate": 2e-05, + "loss": 0.0216, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.627041220664978, + "learning_rate": 2e-05, + "loss": 0.1456, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.3795100450515747, + "learning_rate": 2e-05, + "loss": 0.0442, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.057894088327884674, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.12184643000364304, + "learning_rate": 2e-05, + "loss": 0.0033, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.038447119295597076, + "learning_rate": 2e-05, + "loss": 0.0017, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.07424522191286087, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.024218911305069923, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.08186870813369751, + "learning_rate": 2e-05, + "loss": 0.004, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.009126922115683556, + "learning_rate": 2e-05, + "loss": 0.007, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.013194434344768524, + "learning_rate": 2e-05, + "loss": 0.0035, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.05409262329339981, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7281345129013062, + "learning_rate": 2e-05, + "loss": 0.0257, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.006894941441714764, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.02810726873576641, + "learning_rate": 2e-05, + "loss": 0.345, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.04399490728974342, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.009966861456632614, + "learning_rate": 2e-05, + "loss": 0.0148, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.36275601387023926, + "learning_rate": 2e-05, + "loss": 0.2049, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.004030086100101471, + "learning_rate": 2e-05, + "loss": 0.0072, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.056008920073509216, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.04544645920395851, + "learning_rate": 2e-05, + "loss": 0.0065, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0035371175035834312, + "learning_rate": 2e-05, + "loss": 0.0061, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.64030122756958, + "learning_rate": 2e-05, + "loss": 0.2003, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.0038959316443651915, + "learning_rate": 2e-05, + "loss": 0.0121, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.6076579689979553, + "learning_rate": 2e-05, + "loss": 0.0879, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.186123847961426, + "learning_rate": 2e-05, + "loss": 0.4318, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.024531858041882515, + "learning_rate": 2e-05, + "loss": 0.4271, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.0400189645588398, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.5526515245437622, + "learning_rate": 2e-05, + "loss": 0.0383, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5285302658662400.0, + "train_loss": 0.08479124128818512, + "train_runtime": 96.2218, + "train_samples_per_second": 4.157, + "train_steps_per_second": 1.039 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5285302658662400.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..a538f3797f64e000d87b188c6bf75f02794ca02b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8ee09a9b751b52dfaf12d4be5c48996010af8fccef961b92119d2c6e8fe6faa8 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..4ca918b2cbc514d335e76a42a3badfbff28df83e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cbb6fdfca89ed2fe72839ef3b8902c1fdcb64456891e63f83f4da1f7f566e2f9 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..b513e679f35f7f9a516740fb058a692be41e86d5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:24010f9e1f46a69809db3a2b23e5411971baa1ebfc861d01bea334d9efb8a9ab +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0ca9491d2c3b8b754f8cddafd6ccc7345c4d1f00 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:439cd477c86527b725f012c7e551d9d14e433f7d1e2f6204bd4e311765864b7f +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..f56ea0a0b4406d7d4a40962a46009d5a684ff5bc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29beaf22c1e6d7ef84291cd07135d1769b6177e2aa45cd52f3f9bcbd3f87832d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..85b60348dbf11558e4ac1896fca510bb25fc5006 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4107837f01d11827f7d33c8015065a7fca2dcdc465cab58e0de8d1b927f87711 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..db57ea0d24bce8b48b269dea47101ee853584246 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09b8908a195a6328783c3716c4426934b408bb11c94c942c147404ceb971ebdd +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..c40bded46c9add233c97d211ac1f28d628b55150 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76a1eed9b53de6b4270f115dcafa42bf2c68292300d16d7806d5980d3ef77a38 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..8eb5f39837c86ad65a0eb7b3b3ea53b7e1a99610 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/13_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.784420490264893, + "learning_rate": 2e-05, + "loss": 0.2973, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.1515676975250244, + "learning_rate": 2e-05, + "loss": 0.2302, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.5069868564605713, + "learning_rate": 2e-05, + "loss": 0.3913, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.3688634634017944, + "learning_rate": 2e-05, + "loss": 0.0869, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8190671801567078, + "learning_rate": 2e-05, + "loss": 0.1021, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.20605869591236115, + "learning_rate": 2e-05, + "loss": 0.032, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.787585496902466, + "learning_rate": 2e-05, + "loss": 0.1136, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 4.514202117919922, + "learning_rate": 2e-05, + "loss": 0.3983, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.400081157684326, + "learning_rate": 2e-05, + "loss": 0.1909, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.6767349243164062, + "learning_rate": 2e-05, + "loss": 1.2187, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 6.607516765594482, + "learning_rate": 2e-05, + "loss": 0.5383, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 7.369804859161377, + "learning_rate": 2e-05, + "loss": 0.3169, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.145848751068115, + "learning_rate": 2e-05, + "loss": 0.9271, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.1911654472351074, + "learning_rate": 2e-05, + "loss": 0.1296, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.7649519443511963, + "learning_rate": 2e-05, + "loss": 0.3112, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.9521218538284302, + "learning_rate": 2e-05, + "loss": 0.1392, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.1466885805130005, + "learning_rate": 2e-05, + "loss": 0.2305, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.268568277359009, + "learning_rate": 2e-05, + "loss": 0.3182, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.5398476123809814, + "learning_rate": 2e-05, + "loss": 0.1084, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.485764503479004, + "learning_rate": 2e-05, + "loss": 0.2615, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.05015593394637108, + "learning_rate": 2e-05, + "loss": 0.2594, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.9076626300811768, + "learning_rate": 2e-05, + "loss": 0.3439, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.023381849750876427, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.964611530303955, + "learning_rate": 2e-05, + "loss": 0.2185, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.36583977937698364, + "learning_rate": 2e-05, + "loss": 0.2178, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.364133358001709, + "learning_rate": 2e-05, + "loss": 0.4153, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.9149279594421387, + "learning_rate": 2e-05, + "loss": 0.3808, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.6203057765960693, + "learning_rate": 2e-05, + "loss": 0.1909, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.09775636345148087, + "learning_rate": 2e-05, + "loss": 0.0336, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.12829267978668213, + "learning_rate": 2e-05, + "loss": 0.0748, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 4.191639423370361, + "learning_rate": 2e-05, + "loss": 0.2004, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.5464124083518982, + "learning_rate": 2e-05, + "loss": 0.0699, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.487668752670288, + "learning_rate": 2e-05, + "loss": 0.3548, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.961925506591797, + "learning_rate": 2e-05, + "loss": 0.2961, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.7191128730773926, + "learning_rate": 2e-05, + "loss": 0.2406, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.44289353489875793, + "learning_rate": 2e-05, + "loss": 0.034, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.9182948470115662, + "learning_rate": 2e-05, + "loss": 0.5406, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.6706265211105347, + "learning_rate": 2e-05, + "loss": 0.1153, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.6376776695251465, + "learning_rate": 2e-05, + "loss": 0.2468, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.65388822555542, + "learning_rate": 2e-05, + "loss": 0.089, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.9042004942893982, + "learning_rate": 2e-05, + "loss": 0.1136, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 6.31694221496582, + "learning_rate": 2e-05, + "loss": 0.7479, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 6.465168476104736, + "learning_rate": 2e-05, + "loss": 0.6427, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.535103678703308, + "learning_rate": 2e-05, + "loss": 0.1286, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.845750331878662, + "learning_rate": 2e-05, + "loss": 0.1053, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.446483850479126, + "learning_rate": 2e-05, + "loss": 0.575, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.9901806712150574, + "learning_rate": 2e-05, + "loss": 0.0245, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.963780641555786, + "learning_rate": 2e-05, + "loss": 0.2334, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.4175061285495758, + "learning_rate": 2e-05, + "loss": 0.3401, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.327207565307617, + "learning_rate": 2e-05, + "loss": 0.3079, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5322110830379008.0, + "train_loss": 0.2779290866851807, + "train_runtime": 95.0654, + "train_samples_per_second": 4.208, + "train_steps_per_second": 1.052 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5322110830379008.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..7672ad3d51162d5cbf831fd7593f99be42b2cbe6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8f216033cbe76adbaca648deaf517e95c6e563f94dc7d5419f66a25f4e6ee9d0 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..201766782f6b35f0e70085dd360be04a6ff867f6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1edc36f3dc90960b7c1e5c43c4c8122386fa7878017b310778c118dbcec266b0 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..03637af9dd3f411ef22aeee2736f3093422225cb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:af4b20953b534a89faf0589d05026a86b0cf550c5680145366d22396974899d7 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..c8755e7148f3d0f4ad3b7f6772f280cf0a14ad6c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:df7e4abe95b9e441c7435c4d1cec7f8a727ec0a44f1d1456778ba645913cbab7 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..6d3d59c41e2a84aeb3c8c7545ab8b19c791c96b2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0bd959e08d973aa7d9c60311e5ead8d8a0f92bcfa3a58fd5b6957396dd79af5d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..df9e2c62b07637c1983537de4bfc5f30ec7c5d20 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e684df1819c8d756ec724d1e167cf32907af0409c5c2bb685bca2f71ebfd686c +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..5b7e302e558698268076241baff586ea05770ea1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bc921d2c973abadfa02c6eabdf070229fea2b0eaa47d01bc1f4cf40d952eba8a +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..5778e06fb884418b732d5988f7cf72a59fb29ba4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e4949e2c414d4ab9e33d8ba49632181a61e40718e2203df8026bac574327b810 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..1e9ed2ef7846cb8c5cf3341161d5e09e3c2e8272 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/14_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.6742417812347412, + "learning_rate": 2e-05, + "loss": 0.4682, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.5656232833862305, + "learning_rate": 2e-05, + "loss": 0.2932, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.24041946232318878, + "learning_rate": 2e-05, + "loss": 0.061, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.038180604577064514, + "learning_rate": 2e-05, + "loss": 0.0014, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.618332862854004, + "learning_rate": 2e-05, + "loss": 0.5952, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.793912410736084, + "learning_rate": 2e-05, + "loss": 0.2444, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.261258602142334, + "learning_rate": 2e-05, + "loss": 0.3461, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.025627775117754936, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.2191150188446045, + "learning_rate": 2e-05, + "loss": 0.1418, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.3986196517944336, + "learning_rate": 2e-05, + "loss": 0.0273, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.18143802881240845, + "learning_rate": 2e-05, + "loss": 0.0274, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.02398822084069252, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.950343132019043, + "learning_rate": 2e-05, + "loss": 0.6194, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 7.19309663772583, + "learning_rate": 2e-05, + "loss": 0.346, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.557678699493408, + "learning_rate": 2e-05, + "loss": 0.2076, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.7505941987037659, + "learning_rate": 2e-05, + "loss": 0.0621, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.14495675265789032, + "learning_rate": 2e-05, + "loss": 0.1876, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.7338212728500366, + "learning_rate": 2e-05, + "loss": 0.0411, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.13879674673080444, + "learning_rate": 2e-05, + "loss": 0.0139, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.7909421920776367, + "learning_rate": 2e-05, + "loss": 0.3789, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.22654017806053162, + "learning_rate": 2e-05, + "loss": 0.1435, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.5380771160125732, + "learning_rate": 2e-05, + "loss": 0.3427, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.46424487233161926, + "learning_rate": 2e-05, + "loss": 0.0307, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.036891222000122, + "learning_rate": 2e-05, + "loss": 0.0497, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.029721735045313835, + "learning_rate": 2e-05, + "loss": 0.1723, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.028245823457837105, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.2172374725341797, + "learning_rate": 2e-05, + "loss": 0.0403, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.7566039562225342, + "learning_rate": 2e-05, + "loss": 0.1065, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.7815712690353394, + "learning_rate": 2e-05, + "loss": 0.1853, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.02140204608440399, + "learning_rate": 2e-05, + "loss": 0.0026, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7914906740188599, + "learning_rate": 2e-05, + "loss": 0.1925, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.5331587791442871, + "learning_rate": 2e-05, + "loss": 0.0294, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.9520975351333618, + "learning_rate": 2e-05, + "loss": 0.1145, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.989152431488037, + "learning_rate": 2e-05, + "loss": 0.2315, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.3461293578147888, + "learning_rate": 2e-05, + "loss": 0.172, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.11773144453763962, + "learning_rate": 2e-05, + "loss": 0.025, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.3146727085113525, + "learning_rate": 2e-05, + "loss": 0.1184, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.5638701915740967, + "learning_rate": 2e-05, + "loss": 0.1467, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.8144355416297913, + "learning_rate": 2e-05, + "loss": 0.1162, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.9933357238769531, + "learning_rate": 2e-05, + "loss": 0.0633, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.4231152832508087, + "learning_rate": 2e-05, + "loss": 0.0504, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.6975676417350769, + "learning_rate": 2e-05, + "loss": 0.0589, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.9961870908737183, + "learning_rate": 2e-05, + "loss": 0.0292, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.7348883152008057, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.559696078300476, + "learning_rate": 2e-05, + "loss": 0.0637, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.6239986419677734, + "learning_rate": 2e-05, + "loss": 0.0376, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.14054931700229645, + "learning_rate": 2e-05, + "loss": 0.0166, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.156235694885254, + "learning_rate": 2e-05, + "loss": 0.0479, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.23400920629501343, + "learning_rate": 2e-05, + "loss": 0.0164, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.005347547587007284, + "learning_rate": 2e-05, + "loss": 0.0324, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5337401710870528.0, + "train_loss": 0.13476066112518312, + "train_runtime": 107.6854, + "train_samples_per_second": 3.715, + "train_steps_per_second": 0.929 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5337401710870528.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..2634e962323eb4be7037ab9c15b73f8e0e3a4bd7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:79cced7475eb43b92ae272df94ab972fc6bcbd1e4b9560a7330257797ef8f235 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..2878d5ab0ce04e9848a7dd0794b936dab2d9b485 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f110bb55ead14b3a0084b1c58f6eefd38b1ae1891ee91318aeb36b362d20f041 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..c6207b06223bfa480d1f5ac2e751a3d8a80fc239 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:83434f0c414989133e6a93228dec657c7dcc84abe6bd303882711192a1f05b53 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d8424e209b079e82ca7857224d5f1894f0087f3a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb5e71cfbb6d6585fad21a2108b60c2c9f045672b1365fdfb5c2729655bcc618 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..f00183e1998991591837c92e0a5a3bb7237fecef --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3eeb44ef037b09de6ea8db7735338a3b5cfdfb4ee9af58c68ae492214a225f6d +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..ab506aed344b9a952fea4449b6dfdcd5632dd2c0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b603b9c21520464d97daa72b47b0bd60e1eea1cce30f470afda39d80fc646df +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..34cae1c4cf1a62c0b049d62ae21411066b5b9185 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:204beed1f3b7c7012db4667564448476314d9039dd35d234b89fbdcfeb7d7332 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..dcbdcc84c187df4b296a37f7ac1ec23efdc65ec6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b865a0ee40d51a1f63344121ae3b27b1695c7700e32b1be4456cae3e7ffe65ad +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c12268613b20da6dacf02afb228e0ac548165dd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/15_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.253232717514038, + "learning_rate": 2e-05, + "loss": 0.1395, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 7.638111591339111, + "learning_rate": 2e-05, + "loss": 0.5067, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 5.693814754486084, + "learning_rate": 2e-05, + "loss": 0.4975, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.3362959623336792, + "learning_rate": 2e-05, + "loss": 0.0977, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.7863439917564392, + "learning_rate": 2e-05, + "loss": 0.0803, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.9151372313499451, + "learning_rate": 2e-05, + "loss": 0.1676, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.960108757019043, + "learning_rate": 2e-05, + "loss": 0.0422, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.5828012824058533, + "learning_rate": 2e-05, + "loss": 0.3656, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.9263850450515747, + "learning_rate": 2e-05, + "loss": 0.0725, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.567744016647339, + "learning_rate": 2e-05, + "loss": 0.2089, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 6.0792059898376465, + "learning_rate": 2e-05, + "loss": 0.4536, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 4.688509464263916, + "learning_rate": 2e-05, + "loss": 0.2584, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.132392406463623, + "learning_rate": 2e-05, + "loss": 0.3749, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.3030917644500732, + "learning_rate": 2e-05, + "loss": 0.2617, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.2762231826782227, + "learning_rate": 2e-05, + "loss": 0.1934, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.6769216060638428, + "learning_rate": 2e-05, + "loss": 0.0788, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.9326157569885254, + "learning_rate": 2e-05, + "loss": 0.4451, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.21790820360183716, + "learning_rate": 2e-05, + "loss": 0.0531, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.404603004455566, + "learning_rate": 2e-05, + "loss": 0.3138, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.3446178436279297, + "learning_rate": 2e-05, + "loss": 0.138, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.25932231545448303, + "learning_rate": 2e-05, + "loss": 0.0623, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.951843500137329, + "learning_rate": 2e-05, + "loss": 0.498, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.5786809921264648, + "learning_rate": 2e-05, + "loss": 0.0438, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.585265636444092, + "learning_rate": 2e-05, + "loss": 0.0695, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 5.394436836242676, + "learning_rate": 2e-05, + "loss": 0.6934, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.7639589309692383, + "learning_rate": 2e-05, + "loss": 0.2885, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.15652506053447723, + "learning_rate": 2e-05, + "loss": 0.0146, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.744207501411438, + "learning_rate": 2e-05, + "loss": 0.0315, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.968960762023926, + "learning_rate": 2e-05, + "loss": 0.2654, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.9074415564537048, + "learning_rate": 2e-05, + "loss": 0.0869, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 10.166720390319824, + "learning_rate": 2e-05, + "loss": 0.5512, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 5.295556545257568, + "learning_rate": 2e-05, + "loss": 0.1548, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.9783980846405029, + "learning_rate": 2e-05, + "loss": 0.1093, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 7.5670671463012695, + "learning_rate": 2e-05, + "loss": 0.4659, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.8457675576210022, + "learning_rate": 2e-05, + "loss": 0.4242, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 3.647142171859741, + "learning_rate": 2e-05, + "loss": 0.2458, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.7965811491012573, + "learning_rate": 2e-05, + "loss": 0.2011, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.009468078613281, + "learning_rate": 2e-05, + "loss": 0.3578, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.9808224439620972, + "learning_rate": 2e-05, + "loss": 0.0385, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.350201964378357, + "learning_rate": 2e-05, + "loss": 0.1652, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.088552713394165, + "learning_rate": 2e-05, + "loss": 0.1705, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.107930898666382, + "learning_rate": 2e-05, + "loss": 0.1744, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.167111873626709, + "learning_rate": 2e-05, + "loss": 0.6056, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 7.413355350494385, + "learning_rate": 2e-05, + "loss": 0.7583, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.679892539978027, + "learning_rate": 2e-05, + "loss": 0.2771, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.6088745594024658, + "learning_rate": 2e-05, + "loss": 0.2153, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.7870895862579346, + "learning_rate": 2e-05, + "loss": 0.0673, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.49878594279289246, + "learning_rate": 2e-05, + "loss": 0.8072, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.221227645874023, + "learning_rate": 2e-05, + "loss": 0.2284, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.9601224660873413, + "learning_rate": 2e-05, + "loss": 0.0887, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2213537237696512.0, + "train_loss": 0.258198127746582, + "train_runtime": 64.3219, + "train_samples_per_second": 6.219, + "train_steps_per_second": 1.555 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2213537237696512.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..12fedbb9b7ad5eaa5bda88157531cd2046f4a903 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3c05b84930ea1843041076ea32cc71af1b4886f2af2ea73e3eed59b4969b3326 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..2509311e174af39774b2a4e2917a5811e792931c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f83e4495ad099072a839a8ec34d3f96cb7d981ad5611e1afd6184bfdb58291a6 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a9543c1c9b6a4051f21f76514b36d8a97df38ab0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fc2ec4d50bd1763f364736701604eae1ab1cc5344d826fc75e3c69d573c8cb6a +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..7741501a826dc3b1e7b154343fb5e70e39094a20 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:632c3505767ce0c3ad83d0e00899ae19d1aa3a0042130fd62b0ce8bb9cb523c6 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..d4a05661794dabaccce3feed3a46c2e97a534df2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:904283c780cef21808e02cb192c5d9cd1048256e9a9929b514d971259ee2fb5f +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..9527a8010d8f4a176ecd7c8715a82ba4a6ee2eef --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b673463e311fa4b77932d04fd91ee3c778c539e807cc5b6503f47377e0918360 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c9fe2a6a52a88590b17b5043839821613ac27c70 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0b5df0ec46bb375e2f5bdd6c125d2233049de9eaef27a4b83c25de5e378d3e8b +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..7142805048164a389ea0b5e552654a4fcef6c386 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:31d917630691ee2fa29643b821b3af8bd722f5354fac31ab4ec3592031e1d029 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..011a76007b137ee97c9b7d2a6da31bc747b93212 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/16_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 5.250613689422607, + "learning_rate": 2e-05, + "loss": 0.2937, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.25627562403678894, + "learning_rate": 2e-05, + "loss": 0.0322, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 11.261725425720215, + "learning_rate": 2e-05, + "loss": 0.8574, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 5.502643585205078, + "learning_rate": 2e-05, + "loss": 0.2245, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.149960041046143, + "learning_rate": 2e-05, + "loss": 0.2471, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.11498262733221054, + "learning_rate": 2e-05, + "loss": 0.0312, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.253682851791382, + "learning_rate": 2e-05, + "loss": 0.2339, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.4269609451293945, + "learning_rate": 2e-05, + "loss": 0.3275, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.041144371032715, + "learning_rate": 2e-05, + "loss": 0.0614, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.1361035257577896, + "learning_rate": 2e-05, + "loss": 0.0228, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 6.029748439788818, + "learning_rate": 2e-05, + "loss": 0.6076, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.5250136256217957, + "learning_rate": 2e-05, + "loss": 0.1659, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.7715108394622803, + "learning_rate": 2e-05, + "loss": 0.4926, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.2177114486694336, + "learning_rate": 2e-05, + "loss": 0.1713, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.6275436878204346, + "learning_rate": 2e-05, + "loss": 0.5492, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 4.872981548309326, + "learning_rate": 2e-05, + "loss": 0.2568, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.43277695775032043, + "learning_rate": 2e-05, + "loss": 0.218, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.400402545928955, + "learning_rate": 2e-05, + "loss": 0.2136, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.6360125541687012, + "learning_rate": 2e-05, + "loss": 0.0765, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.8332691192626953, + "learning_rate": 2e-05, + "loss": 0.2413, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.12378638237714767, + "learning_rate": 2e-05, + "loss": 0.007, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7434951663017273, + "learning_rate": 2e-05, + "loss": 0.1788, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.385206699371338, + "learning_rate": 2e-05, + "loss": 0.0743, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.5788016319274902, + "learning_rate": 2e-05, + "loss": 0.0722, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.2752498686313629, + "learning_rate": 2e-05, + "loss": 0.048, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.455329179763794, + "learning_rate": 2e-05, + "loss": 0.0341, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.206432580947876, + "learning_rate": 2e-05, + "loss": 0.0483, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 7.318122863769531, + "learning_rate": 2e-05, + "loss": 0.3824, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.380479574203491, + "learning_rate": 2e-05, + "loss": 0.2392, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.404880404472351, + "learning_rate": 2e-05, + "loss": 0.0618, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.5300942063331604, + "learning_rate": 2e-05, + "loss": 0.0182, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.9066197872161865, + "learning_rate": 2e-05, + "loss": 0.561, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.7215485572814941, + "learning_rate": 2e-05, + "loss": 0.0793, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 10.421614646911621, + "learning_rate": 2e-05, + "loss": 1.5366, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.35321307182312, + "learning_rate": 2e-05, + "loss": 0.2069, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 9.214303970336914, + "learning_rate": 2e-05, + "loss": 0.4011, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 6.32353401184082, + "learning_rate": 2e-05, + "loss": 0.7807, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.07320781797170639, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.2380521297454834, + "learning_rate": 2e-05, + "loss": 0.0374, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.3714147210121155, + "learning_rate": 2e-05, + "loss": 0.0312, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 5.273041248321533, + "learning_rate": 2e-05, + "loss": 0.511, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.787007212638855, + "learning_rate": 2e-05, + "loss": 0.1908, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.277595043182373, + "learning_rate": 2e-05, + "loss": 0.024, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.134889125823975, + "learning_rate": 2e-05, + "loss": 0.2282, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.9650557041168213, + "learning_rate": 2e-05, + "loss": 0.356, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.3101871013641357, + "learning_rate": 2e-05, + "loss": 0.1061, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 7.785775661468506, + "learning_rate": 2e-05, + "loss": 0.7836, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 5.7999491691589355, + "learning_rate": 2e-05, + "loss": 0.1265, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.1724660396575928, + "learning_rate": 2e-05, + "loss": 0.0916, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.3783414363861084, + "learning_rate": 2e-05, + "loss": 0.1393, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206762417520640.0, + "train_loss": 0.25367055892944335, + "train_runtime": 63.6649, + "train_samples_per_second": 6.283, + "train_steps_per_second": 1.571 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206762417520640.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..88b99dc146ab87b0d59b0e009755f91871c508f9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4543d95014adfc2b90d2c6354e9387910e59d896421ef4b053b920c96274167a +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..d396e1b928b477077645489e7d031eec2204c123 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:33db30bacac93e84d224b9de950d5ad8f31fd2ee4ee36ecacab23822b729928c +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..de05148d8d0f3bfe1a10a21d24e8a0416fe25722 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e8cc8afd3f36677873b6c1fdb8dcfc34f95d9d92d423916302757598239c848f +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b5f762143200834e66ad9c1beb3b650dd9d138e0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:02c7e632b91e93a1c5b2040ec59c66098ee3359cae5fab14090b1fcb4a659d1b +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..03244ec31f1e8ed85d92ef15b01dc745ed5d3283 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4731163b732c51f7c73543f7df26d020fc83b484e75376cbd22e61ebc46f62c8 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..0ed1718d7d69c53e2f96b09d47a83d1f4ee07529 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:100e067c752f4d26c9af6dfe74499f1fb4c1f827d6ec5d51f14aeafc9e6faa3b +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..dd3e26bdf76e3de83cb84ab2964b63b5896facba --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c9cd72f8145fe458c9a484586611a6e8d22390005881f449074e432c904b8f5f +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..4e2704860f4ed919d418b8dae9428c6875f87c4e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ab063d80174d9240106308cfe34c4a6c75fbbfb63896e9a3127e52e288ea3504 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..b562253199abb668c0693a68b1e08e9ea6390927 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/17_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.4537246227264404, + "learning_rate": 2e-05, + "loss": 0.0612, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.458911418914795, + "learning_rate": 2e-05, + "loss": 0.4085, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.7891697883605957, + "learning_rate": 2e-05, + "loss": 0.1109, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.116798996925354, + "learning_rate": 2e-05, + "loss": 0.0561, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.723391056060791, + "learning_rate": 2e-05, + "loss": 0.1316, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.115161180496216, + "learning_rate": 2e-05, + "loss": 0.1964, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.9211235046386719, + "learning_rate": 2e-05, + "loss": 0.0272, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 6.994407653808594, + "learning_rate": 2e-05, + "loss": 0.2409, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 7.782790660858154, + "learning_rate": 2e-05, + "loss": 0.3343, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 8.48725700378418, + "learning_rate": 2e-05, + "loss": 0.9732, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.1843607425689697, + "learning_rate": 2e-05, + "loss": 0.1797, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 8.04741382598877, + "learning_rate": 2e-05, + "loss": 0.7436, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.008490800857544, + "learning_rate": 2e-05, + "loss": 0.0914, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.12534895539283752, + "learning_rate": 2e-05, + "loss": 0.0072, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.5183145999908447, + "learning_rate": 2e-05, + "loss": 0.1687, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.4072463512420654, + "learning_rate": 2e-05, + "loss": 0.093, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.389726161956787, + "learning_rate": 2e-05, + "loss": 0.0672, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.5420376062393188, + "learning_rate": 2e-05, + "loss": 0.0215, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.107137441635132, + "learning_rate": 2e-05, + "loss": 0.1958, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.8174989223480225, + "learning_rate": 2e-05, + "loss": 0.2697, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.6306998133659363, + "learning_rate": 2e-05, + "loss": 0.0259, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.2154836654663086, + "learning_rate": 2e-05, + "loss": 0.3511, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.9298397898674011, + "learning_rate": 2e-05, + "loss": 0.0343, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 4.799689292907715, + "learning_rate": 2e-05, + "loss": 0.2344, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.4766454696655273, + "learning_rate": 2e-05, + "loss": 0.3545, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.9954042434692383, + "learning_rate": 2e-05, + "loss": 0.2381, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.259938955307007, + "learning_rate": 2e-05, + "loss": 0.2236, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 5.574470520019531, + "learning_rate": 2e-05, + "loss": 0.2836, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 7.461014270782471, + "learning_rate": 2e-05, + "loss": 0.3919, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 8.046769142150879, + "learning_rate": 2e-05, + "loss": 0.4047, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 9.566935539245605, + "learning_rate": 2e-05, + "loss": 1.0526, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 5.701999664306641, + "learning_rate": 2e-05, + "loss": 0.9368, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.5290571451187134, + "learning_rate": 2e-05, + "loss": 0.0701, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 7.140415668487549, + "learning_rate": 2e-05, + "loss": 0.2431, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.44766056537628174, + "learning_rate": 2e-05, + "loss": 0.035, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 4.563236236572266, + "learning_rate": 2e-05, + "loss": 0.4814, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.113389492034912, + "learning_rate": 2e-05, + "loss": 0.3503, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.00366735458374, + "learning_rate": 2e-05, + "loss": 0.6964, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 3.752512216567993, + "learning_rate": 2e-05, + "loss": 0.2393, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.3487975299358368, + "learning_rate": 2e-05, + "loss": 0.1346, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.205235481262207, + "learning_rate": 2e-05, + "loss": 0.1815, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.5800580978393555, + "learning_rate": 2e-05, + "loss": 0.3462, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.081882476806641, + "learning_rate": 2e-05, + "loss": 0.3831, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.183148741722107, + "learning_rate": 2e-05, + "loss": 0.0339, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 5.5165486335754395, + "learning_rate": 2e-05, + "loss": 0.5569, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.472242832183838, + "learning_rate": 2e-05, + "loss": 0.1401, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.08788532763719559, + "learning_rate": 2e-05, + "loss": 0.0341, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.0250599384307861, + "learning_rate": 2e-05, + "loss": 0.057, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.023240989074110985, + "learning_rate": 2e-05, + "loss": 0.0748, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.5438010692596436, + "learning_rate": 2e-05, + "loss": 0.0689, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2211830323740672.0, + "train_loss": 0.26072967052459717, + "train_runtime": 62.1578, + "train_samples_per_second": 6.435, + "train_steps_per_second": 1.609 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2211830323740672.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..ad9c6401217467697c80cef5c091b2bb6687467d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7db1a0127f01d59914a8dd9e1940e0f25a0e9754d7dc4623784359fa7278581e +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..cf710cdd89fdfdda35820fd807f230c83d3f3a70 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e6e1d1b00b1b8469ce7641b54faec0c1bdc8155e4f69ee1a39d51aef9db48034 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..08eae2b2d26d1e3eb95f6ee1b63c6667ebdb9e19 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9de25d7374cf5e28dbab12fb81bc57d85b4ef9346ead28331236e85351161799 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..36997758f3c567eea967d17c3ec926c96b754bbc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e76e237122dc4d1e5938b545244a5bf3ee2f9f709a391ca7916ff6ead48a2e1c +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..b1b3a1e74f4a47e2dbdc627938b554300151ddce --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bbaf66519ce1d9245d74abdf779348f973be9a4ef47f4e987bee9cf993e8cee9 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b958e720d49aeb67c893d151834456a37b502863 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bef5ffdddbcb3be4aab808a1ad4ea9c1a49adf379cc993997d664901e4b97858 +size 368444338 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..0e0228b88dff1f824da8cb7d2c255a170c653973 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:169a3777d70711616d1e94207229747b9619cda22b5f1d9de0ef76b5d28db0ba +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..b198b588c57736ec6ae51f713f94238f5ca15cc1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:13fc27ae54390ad16d501ac6b0517fe675c132e0bc50f1bb2f70bfb3411fad9c +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..9d5096d5a77378926e0fe950a16b9c77cd8f0bd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/18_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.11914084106683731, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.0264544486999512, + "learning_rate": 2e-05, + "loss": 0.0554, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.48574885725975037, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.07401999086141586, + "learning_rate": 2e-05, + "loss": 0.0025, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.07262499630451202, + "learning_rate": 2e-05, + "loss": 0.4016, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 6.500160217285156, + "learning_rate": 2e-05, + "loss": 0.1585, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.213310956954956, + "learning_rate": 2e-05, + "loss": 0.0967, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.3930889368057251, + "learning_rate": 2e-05, + "loss": 0.0414, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.391648769378662, + "learning_rate": 2e-05, + "loss": 0.0434, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.271697759628296, + "learning_rate": 2e-05, + "loss": 0.1466, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.13101613521575928, + "learning_rate": 2e-05, + "loss": 0.1262, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.37160950899124146, + "learning_rate": 2e-05, + "loss": 0.1028, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.42677852511405945, + "learning_rate": 2e-05, + "loss": 0.0705, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 5.074921131134033, + "learning_rate": 2e-05, + "loss": 0.4593, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.06918313354253769, + "learning_rate": 2e-05, + "loss": 0.0676, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 5.970218658447266, + "learning_rate": 2e-05, + "loss": 0.1877, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.7115309834480286, + "learning_rate": 2e-05, + "loss": 0.0331, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.21321725845336914, + "learning_rate": 2e-05, + "loss": 0.1753, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.9680980443954468, + "learning_rate": 2e-05, + "loss": 0.084, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.253144770860672, + "learning_rate": 2e-05, + "loss": 0.0242, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.2706148624420166, + "learning_rate": 2e-05, + "loss": 0.306, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.20665663480758667, + "learning_rate": 2e-05, + "loss": 0.0073, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.9926193952560425, + "learning_rate": 2e-05, + "loss": 0.3556, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.1926522254943848, + "learning_rate": 2e-05, + "loss": 0.1227, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.2155959606170654, + "learning_rate": 2e-05, + "loss": 0.0785, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.424222469329834, + "learning_rate": 2e-05, + "loss": 0.1125, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 4.744777679443359, + "learning_rate": 2e-05, + "loss": 0.2157, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.2267092615365982, + "learning_rate": 2e-05, + "loss": 0.3012, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.22740937769412994, + "learning_rate": 2e-05, + "loss": 0.0821, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.9129446744918823, + "learning_rate": 2e-05, + "loss": 0.0802, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.008006921038031578, + "learning_rate": 2e-05, + "loss": 0.031, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.9925302863121033, + "learning_rate": 2e-05, + "loss": 0.0636, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.02182718925178051, + "learning_rate": 2e-05, + "loss": 0.1588, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.4735757112503052, + "learning_rate": 2e-05, + "loss": 0.0742, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 6.51339054107666, + "learning_rate": 2e-05, + "loss": 0.4195, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.06425304710865021, + "learning_rate": 2e-05, + "loss": 0.0081, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.159090518951416, + "learning_rate": 2e-05, + "loss": 0.1075, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 9.029239654541016, + "learning_rate": 2e-05, + "loss": 0.2104, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 5.890206336975098, + "learning_rate": 2e-05, + "loss": 0.2701, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.4602356553077698, + "learning_rate": 2e-05, + "loss": 0.0113, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 5.5746235847473145, + "learning_rate": 2e-05, + "loss": 0.2881, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0744781345129013, + "learning_rate": 2e-05, + "loss": 0.1524, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.110683798789978, + "learning_rate": 2e-05, + "loss": 0.0202, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.09108743816614151, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.776930570602417, + "learning_rate": 2e-05, + "loss": 0.1779, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.8053197860717773, + "learning_rate": 2e-05, + "loss": 0.0941, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 6.740932941436768, + "learning_rate": 2e-05, + "loss": 0.3666, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 4.950959205627441, + "learning_rate": 2e-05, + "loss": 0.0771, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.036170922219753265, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 6.338098049163818, + "learning_rate": 2e-05, + "loss": 0.2845, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206568850391040.0, + "train_loss": 0.1356373357772827, + "train_runtime": 76.6897, + "train_samples_per_second": 5.216, + "train_steps_per_second": 1.304 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206568850391040.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..16cc67c7f07798c69240b6795cf091db582c4c35 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e582fdefebe6058190f8d785d8b760d857f33763a51ffe3780d7cbc106dd8419 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..6180b85f820863d9f0ea88d7d4f6fbfaca05491a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63f1811a868265cb70c1ce40289be23d275e6296593bcbfc882b4145b313c830 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..396c7dee3bb3aa9c91c325b1255a38f03128a992 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8324ecfa23506264f54ef27eb97da0d9633394956de40709552c25b19aede21f +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..14ff8ba5aeff134a97d7438515ad6b127d7d81ef --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf853c7bd8062b3e9f23a3b9aad068e15a14dcdfcf1eadf2e0031f468af36f19 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..da6a4258d354881eb32cd0b56f99c0b6892ee9a6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5ecde1d2c67243e31e4d434ec432f737b6f218adb33be9a5bb2e74e455af24e5 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..65ee8d9edc2bcee6b07b1c0917abda56f9fdcb7c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97a6272d751f7ec38d94b210933b6f434702560bf33ea61ec39ea80b045a454f +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..4babe107b3be9af1f1595487e8b0d708d21dd99b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c14ea2278efc3f1629fb9edd559b114379847501d34b5d7e141388cb3cee1b1 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..01bdd747da809ba2be67f007a4f1f4b716fe726d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ce4a4f92cafae066f91db9c4951a06d170145f6e86671be18d3869e384da40d4 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..f82699c7b70abbce32fb9219ff3bd5a5a50c3d35 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/19_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.5865542888641357, + "learning_rate": 2e-05, + "loss": 0.0872, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.40844088792800903, + "learning_rate": 2e-05, + "loss": 0.0836, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.8873893618583679, + "learning_rate": 2e-05, + "loss": 0.0586, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.4208688735961914, + "learning_rate": 2e-05, + "loss": 0.1641, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.047400951385498, + "learning_rate": 2e-05, + "loss": 0.4212, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.4284755289554596, + "learning_rate": 2e-05, + "loss": 0.0495, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.143988013267517, + "learning_rate": 2e-05, + "loss": 0.0777, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.781654417514801, + "learning_rate": 2e-05, + "loss": 0.1336, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.4147431552410126, + "learning_rate": 2e-05, + "loss": 0.2243, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.17606163024902344, + "learning_rate": 2e-05, + "loss": 0.0133, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.2263504266738892, + "learning_rate": 2e-05, + "loss": 0.0946, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.2530275583267212, + "learning_rate": 2e-05, + "loss": 0.0548, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.1259047985076904, + "learning_rate": 2e-05, + "loss": 0.1152, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.9004509449005127, + "learning_rate": 2e-05, + "loss": 0.1245, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.063126564025879, + "learning_rate": 2e-05, + "loss": 0.2131, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.5045535564422607, + "learning_rate": 2e-05, + "loss": 0.4827, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.3128724098205566, + "learning_rate": 2e-05, + "loss": 0.2193, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.4230955243110657, + "learning_rate": 2e-05, + "loss": 0.0839, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2731386721134186, + "learning_rate": 2e-05, + "loss": 0.0383, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.10574228316545486, + "learning_rate": 2e-05, + "loss": 0.0553, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.30341637134552, + "learning_rate": 2e-05, + "loss": 0.0691, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7762811779975891, + "learning_rate": 2e-05, + "loss": 0.5232, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.8278684616088867, + "learning_rate": 2e-05, + "loss": 0.1438, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.4772418141365051, + "learning_rate": 2e-05, + "loss": 0.2842, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.29748237133026123, + "learning_rate": 2e-05, + "loss": 0.2724, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.043168544769287, + "learning_rate": 2e-05, + "loss": 0.2463, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.29596468806266785, + "learning_rate": 2e-05, + "loss": 0.1825, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.12161802500486374, + "learning_rate": 2e-05, + "loss": 0.0147, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.438669443130493, + "learning_rate": 2e-05, + "loss": 0.1407, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.7590668201446533, + "learning_rate": 2e-05, + "loss": 0.08, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.1044471338391304, + "learning_rate": 2e-05, + "loss": 0.1786, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.5248003005981445, + "learning_rate": 2e-05, + "loss": 0.2643, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.5164191722869873, + "learning_rate": 2e-05, + "loss": 0.284, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.06531926244497299, + "learning_rate": 2e-05, + "loss": 0.0775, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.5524253845214844, + "learning_rate": 2e-05, + "loss": 0.4005, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.4818257689476013, + "learning_rate": 2e-05, + "loss": 0.1743, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.25305014848709106, + "learning_rate": 2e-05, + "loss": 0.0277, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 7.797881603240967, + "learning_rate": 2e-05, + "loss": 0.5299, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 3.200997829437256, + "learning_rate": 2e-05, + "loss": 0.1512, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.419097423553467, + "learning_rate": 2e-05, + "loss": 0.1496, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.4525766372680664, + "learning_rate": 2e-05, + "loss": 0.3453, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.13708224892616272, + "learning_rate": 2e-05, + "loss": 0.0241, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.002226962009444833, + "learning_rate": 2e-05, + "loss": 0.0459, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.10485492646694183, + "learning_rate": 2e-05, + "loss": 0.3184, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3034449815750122, + "learning_rate": 2e-05, + "loss": 0.0428, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.2066220045089722, + "learning_rate": 2e-05, + "loss": 0.073, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.7168982625007629, + "learning_rate": 2e-05, + "loss": 0.0897, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.0437278747558594, + "learning_rate": 2e-05, + "loss": 0.1707, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.20914383232593536, + "learning_rate": 2e-05, + "loss": 0.0392, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.08821244537830353, + "learning_rate": 2e-05, + "loss": 0.2855, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293519879012352.0, + "train_loss": 0.16848104000091552, + "train_runtime": 93.2063, + "train_samples_per_second": 4.292, + "train_steps_per_second": 1.073 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293519879012352.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..664d4e8c58ab208195325959b67957247728a3d6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fc1b3351373c77fe9f30a792069e790c375067b7b8a53a8b069868c144e58ecc +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..06e913eb547162e16ca668a4d6d214903d91f23b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7306b14c84b1f6fc6f0372e09f6d458a863b753d058241d64ecf1578e3378535 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..616d55b7f4bfc68f90e1fe83db5d7b5e7de127ee --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7b9e6af59547f9eb81ba2dca9eca12cf8de60719aa83eb1652c0aab9ea8a010b +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..ad7d92577d2a101c83c176feb1803f6baa6866ac --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2883b40042646f46d2399f29ad5405bdfb3c63ba76524dad9e301b51f3ef88f6 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..bdb61cfae7cd87fea47eda54e971f80ca64c3585 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:092267629c20a0d44b634c39faec128b23949c289f083ad2c264d3835696a3ff +size 368442474 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..e08ac2c7611918eb555c23ad1ddc719f629588de --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:385245e2e0c537e551b957a19d00f854396634b9491febb3c4061ad68350b6f6 +size 368443438 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..bf512c0339de98987359909b69a0a1657c33430d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7383b12d37cb0e84b1fb057427dc02626aff0e59faa0cb839acfc3cffd995b78 +size 368442474 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..a144a008e67296c16b734c10159870c47482ec76 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c5d44314420c75adf603052b7809c9517687646c0e202fcb4802349c17ee7496 +size 368442474 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..90f875f7cc5219690766ae3d1186a7115970044f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/1_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.04646485671401024, + "learning_rate": 2e-05, + "loss": 0.0111, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.046799346804618835, + "learning_rate": 2e-05, + "loss": 0.0173, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.12124989926815033, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.1149299144744873, + "learning_rate": 2e-05, + "loss": 0.0057, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.0843510627746582, + "learning_rate": 2e-05, + "loss": 0.0017, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.6625175476074219, + "learning_rate": 2e-05, + "loss": 0.0348, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.019028062000870705, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.09930849075317383, + "learning_rate": 2e-05, + "loss": 0.0574, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.06192668154835701, + "learning_rate": 2e-05, + "loss": 0.0494, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.005179049912840128, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.025291118770837784, + "learning_rate": 2e-05, + "loss": 0.024, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.08431576937437057, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.673893451690674, + "learning_rate": 2e-05, + "loss": 0.0734, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.024287063628435135, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.15580369532108307, + "learning_rate": 2e-05, + "loss": 0.0035, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.21148237586021423, + "learning_rate": 2e-05, + "loss": 0.0199, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.7502453327178955, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.002698663156479597, + "learning_rate": 2e-05, + "loss": 0.0001, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.03815506026148796, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.9533177018165588, + "learning_rate": 2e-05, + "loss": 0.0174, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.11492998152971268, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.5814619064331055, + "learning_rate": 2e-05, + "loss": 0.1805, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.005168592557311058, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.01909978874027729, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.008409270085394382, + "learning_rate": 2e-05, + "loss": 0.0009, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.030386749655008316, + "learning_rate": 2e-05, + "loss": 0.0213, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.424407720565796, + "learning_rate": 2e-05, + "loss": 0.0491, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.020852098241448402, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.007163367699831724, + "learning_rate": 2e-05, + "loss": 0.0885, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.007970490492880344, + "learning_rate": 2e-05, + "loss": 0.0255, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.21955572068691254, + "learning_rate": 2e-05, + "loss": 0.0025, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.015326420776546001, + "learning_rate": 2e-05, + "loss": 0.0057, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.01760123483836651, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.18366923928260803, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.00444579916074872, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.4639798700809479, + "learning_rate": 2e-05, + "loss": 0.0066, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.8693036437034607, + "learning_rate": 2e-05, + "loss": 0.0161, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.008622714318335056, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.8267555236816406, + "learning_rate": 2e-05, + "loss": 0.032, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.0077345846220850945, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.007618137635290623, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0009710122831165791, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.016963917762041092, + "learning_rate": 2e-05, + "loss": 0.1231, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.013078974559903145, + "learning_rate": 2e-05, + "loss": 0.0017, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.13917967677116394, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.5520843863487244, + "learning_rate": 2e-05, + "loss": 0.0039, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.007892481982707977, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.2892105281352997, + "learning_rate": 2e-05, + "loss": 0.005, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.006346290931105614, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.03375209867954254, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2217003838341120.0, + "train_loss": 0.01863267719745636, + "train_runtime": 60.933, + "train_samples_per_second": 6.565, + "train_steps_per_second": 1.641 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2217003838341120.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..28fe5485f190c28c10192550e39263d2e83e8be7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a3935cce0a3cdcf555316a47926148c004ca51640aef26b8ef3cc46c92774877 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..308712adbf914375043b78ae33d4e8d4b9f297bd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4bfaecaeb87a8f6d1b17af52c7d58d2567327293695093b099631b254cd7d133 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..0ad899a3e573d39d28c5d0c6d1318c5d787687db --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56599fd4fe632a8540bc4eb2577a7fad454d39ee216388738de30bf40b859fd3 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0b16613445bd25a1811d9a8a6ee48ec2184f9cc2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:da5850e1f0f9387416fb084323b074b08e8a118ac43370e6a68600f9695862e4 +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..785b62fb6f51ad15ff5b48f04de6b184de503930 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7fdef5f66297f333fd779b49c893e75bf8e5ee5894fe001addd54e89af8f1ac2 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..922f4bd70afe82efff9c4e6e0237056ce57982f4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cb20c8b7633b7cfe30b58c0feb74e0e05ae18359ae43174257d43d4cb860833c +size 791579754 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..fbc4ac01672ee0ba2f0e3a099ce1e34950754064 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bde79d50257856ce3fbd162da9484ce0942704dfb8f278bb8d754370a3f73d05 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..ec3af8fac6f5934d9033c1e2c6d9e2095600ddcf --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c1e3ac6dd70dc0b31e297b5ef05a2e4857786f3a7d68f562d3475550a2611582 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..cada8c471b0384a5ae1012747bd540ec7b6be50e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/20_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.13359451293945312, + "learning_rate": 2e-05, + "loss": 0.0257, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.0686542987823486, + "learning_rate": 2e-05, + "loss": 0.0892, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.891756296157837, + "learning_rate": 2e-05, + "loss": 0.1078, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.989849865436554, + "learning_rate": 2e-05, + "loss": 0.0332, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.10513162612915, + "learning_rate": 2e-05, + "loss": 0.4802, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.357785701751709, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.45844823122024536, + "learning_rate": 2e-05, + "loss": 0.0468, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.373875617980957, + "learning_rate": 2e-05, + "loss": 0.0803, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.6453001499176025, + "learning_rate": 2e-05, + "loss": 0.1974, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.11199987679719925, + "learning_rate": 2e-05, + "loss": 0.135, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.4586637020111084, + "learning_rate": 2e-05, + "loss": 0.0158, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.35699617862701416, + "learning_rate": 2e-05, + "loss": 0.0149, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.1243268251419067, + "learning_rate": 2e-05, + "loss": 0.2418, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.20649310946464539, + "learning_rate": 2e-05, + "loss": 0.0128, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.05033278465271, + "learning_rate": 2e-05, + "loss": 0.0274, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.2613182067871094, + "learning_rate": 2e-05, + "loss": 0.2411, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.12879742681980133, + "learning_rate": 2e-05, + "loss": 0.0441, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.49849948287010193, + "learning_rate": 2e-05, + "loss": 0.0331, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2650819718837738, + "learning_rate": 2e-05, + "loss": 0.0824, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.6074606776237488, + "learning_rate": 2e-05, + "loss": 0.0522, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.38714420795440674, + "learning_rate": 2e-05, + "loss": 0.1827, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.3373780250549316, + "learning_rate": 2e-05, + "loss": 0.076, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.3159501552581787, + "learning_rate": 2e-05, + "loss": 0.1001, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.9144917726516724, + "learning_rate": 2e-05, + "loss": 0.0308, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.5086147785186768, + "learning_rate": 2e-05, + "loss": 0.1402, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.358862042427063, + "learning_rate": 2e-05, + "loss": 0.0117, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.8163649439811707, + "learning_rate": 2e-05, + "loss": 0.1331, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.06936631351709366, + "learning_rate": 2e-05, + "loss": 0.0378, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.011788858100771904, + "learning_rate": 2e-05, + "loss": 0.1569, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.1721450835466385, + "learning_rate": 2e-05, + "loss": 0.0529, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.5808074474334717, + "learning_rate": 2e-05, + "loss": 0.243, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 5.087819576263428, + "learning_rate": 2e-05, + "loss": 0.5858, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 6.965622425079346, + "learning_rate": 2e-05, + "loss": 0.8763, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 4.730805397033691, + "learning_rate": 2e-05, + "loss": 0.1691, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.6084457039833069, + "learning_rate": 2e-05, + "loss": 0.0424, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.08248386532068253, + "learning_rate": 2e-05, + "loss": 0.0323, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.11081436276435852, + "learning_rate": 2e-05, + "loss": 0.0109, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 7.164796829223633, + "learning_rate": 2e-05, + "loss": 0.6904, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.4661777913570404, + "learning_rate": 2e-05, + "loss": 0.0324, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.01923215389251709, + "learning_rate": 2e-05, + "loss": 0.0182, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.14167776703834534, + "learning_rate": 2e-05, + "loss": 0.0946, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.9970797896385193, + "learning_rate": 2e-05, + "loss": 0.0891, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.3683152496814728, + "learning_rate": 2e-05, + "loss": 0.0258, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.3326032161712646, + "learning_rate": 2e-05, + "loss": 0.2027, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3553727865219116, + "learning_rate": 2e-05, + "loss": 0.1235, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.134339332580566, + "learning_rate": 2e-05, + "loss": 0.3354, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.7308475375175476, + "learning_rate": 2e-05, + "loss": 0.0422, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.10863185673952103, + "learning_rate": 2e-05, + "loss": 0.004, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.7102224826812744, + "learning_rate": 2e-05, + "loss": 0.1505, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.1631685495376587, + "learning_rate": 2e-05, + "loss": 0.0068, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293985824243712.0, + "train_loss": 0.13342400908470153, + "train_runtime": 100.5195, + "train_samples_per_second": 3.979, + "train_steps_per_second": 0.995 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293985824243712.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f1addc24171c66597bc4f20dc85e2cbee7022337 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3b812963fcd3cf918f1feebf443318c7df7767f179dc194a1e2aaf1ff587a13 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..64a5162acd8892ed17833ade5ed864b2d5eec77f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09ebb7f1cc10e1ef474d8526da94070d2250a7e4dc31487a39ef5d109c3c6a2b +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..2732ab4dca72baeb0bb4a8c17b47569d19a7db9d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:431407e07f6da5d778c2be0a1f404b291c95e0a5345e7cd3c4dd1d0c7467d40c +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..a4185d1a38ca144da23f144004efa6365c912dd0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be57f501ab029219c45dc2e53fb1957ab3955da8777183c505afb9dadf021ade +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..fc2c8433c386eece5d301717306e66fb16147d47 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fa0052ada693663da4af1b4713924dcc8c5b04f1465dbd7b052cb482d50902c +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..44dd37b1855920aae6f743708c7b9a6bb1f57fcd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7eb9f90a24824c4611c6e2a558dc219b2417c2655cdcf4ae8b680f41b17dbda2 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..8e7b2887e3f1800e89bd837a16c5ce58caa23c7e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:64ec709677c57037bd7a4eb0d9993641d87142005c0b9e01f7204e38e5c465e9 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..4d0c02227dbac551274c8566f178507e66fd2e66 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e9eddfe4cf7ad4d6f60bfab46dfdc299f2739b789425fa49b486afe642ea4c41 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..48e75f38e5a1af85c22ee10270c47b9f3d3bf6f0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/2_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.3720414638519287, + "learning_rate": 2e-05, + "loss": 0.1224, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.2211227416992188, + "learning_rate": 2e-05, + "loss": 0.6421, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.12983775138855, + "learning_rate": 2e-05, + "loss": 0.3921, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.81449294090271, + "learning_rate": 2e-05, + "loss": 0.3659, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.8250796794891357, + "learning_rate": 2e-05, + "loss": 0.2006, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.933518409729004, + "learning_rate": 2e-05, + "loss": 0.4839, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.27137210965156555, + "learning_rate": 2e-05, + "loss": 0.1895, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.0695507526397705, + "learning_rate": 2e-05, + "loss": 0.3375, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.9218456745147705, + "learning_rate": 2e-05, + "loss": 0.2461, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.109397292137146, + "learning_rate": 2e-05, + "loss": 0.4651, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.4700283706188202, + "learning_rate": 2e-05, + "loss": 0.0689, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.6726670265197754, + "learning_rate": 2e-05, + "loss": 0.2531, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.765824317932129, + "learning_rate": 2e-05, + "loss": 0.512, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.386554479598999, + "learning_rate": 2e-05, + "loss": 0.1865, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.473822683095932, + "learning_rate": 2e-05, + "loss": 0.0886, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.360491991043091, + "learning_rate": 2e-05, + "loss": 0.2003, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.1749348640441895, + "learning_rate": 2e-05, + "loss": 0.3293, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.0480343103408813, + "learning_rate": 2e-05, + "loss": 0.0514, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.0307161808013916, + "learning_rate": 2e-05, + "loss": 0.5656, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 5.0201897621154785, + "learning_rate": 2e-05, + "loss": 0.3982, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.1526058912277222, + "learning_rate": 2e-05, + "loss": 0.3188, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7723930478096008, + "learning_rate": 2e-05, + "loss": 0.0772, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.877447783946991, + "learning_rate": 2e-05, + "loss": 0.2079, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.354142963886261, + "learning_rate": 2e-05, + "loss": 0.1499, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.2224998474121094, + "learning_rate": 2e-05, + "loss": 0.1144, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.5339128971099854, + "learning_rate": 2e-05, + "loss": 0.292, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 4.064701080322266, + "learning_rate": 2e-05, + "loss": 0.2969, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 4.4953765869140625, + "learning_rate": 2e-05, + "loss": 0.2523, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.2472304105758667, + "learning_rate": 2e-05, + "loss": 0.0538, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.1343146413564682, + "learning_rate": 2e-05, + "loss": 0.1115, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.078665852546692, + "learning_rate": 2e-05, + "loss": 0.3179, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.421689987182617, + "learning_rate": 2e-05, + "loss": 0.1613, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.780685544013977, + "learning_rate": 2e-05, + "loss": 0.095, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.869209051132202, + "learning_rate": 2e-05, + "loss": 0.2301, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.6825551986694336, + "learning_rate": 2e-05, + "loss": 0.3974, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.1026113033294678, + "learning_rate": 2e-05, + "loss": 0.4912, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.000330924987793, + "learning_rate": 2e-05, + "loss": 0.034, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.0241987705230713, + "learning_rate": 2e-05, + "loss": 0.1639, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.7751209735870361, + "learning_rate": 2e-05, + "loss": 0.042, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 6.815292835235596, + "learning_rate": 2e-05, + "loss": 0.5489, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 8.408271789550781, + "learning_rate": 2e-05, + "loss": 1.7144, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.2357848882675171, + "learning_rate": 2e-05, + "loss": 0.0103, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.995177984237671, + "learning_rate": 2e-05, + "loss": 0.2575, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 4.14724063873291, + "learning_rate": 2e-05, + "loss": 0.3335, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.526899814605713, + "learning_rate": 2e-05, + "loss": 0.8911, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.20602548122406006, + "learning_rate": 2e-05, + "loss": 0.1156, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.380793571472168, + "learning_rate": 2e-05, + "loss": 0.2468, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.7478464841842651, + "learning_rate": 2e-05, + "loss": 0.1398, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.965731620788574, + "learning_rate": 2e-05, + "loss": 0.693, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.9956061840057373, + "learning_rate": 2e-05, + "loss": 0.0876, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5221258941693952.0, + "train_loss": 0.2989007043838501, + "train_runtime": 107.0645, + "train_samples_per_second": 3.736, + "train_steps_per_second": 0.934 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5221258941693952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..803fc99ec0173970f5af2d599d83bc7d22472288 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5df47794375faf467d41dc27986cdd3aa382041e64d4b9a61304960cc6fbfe68 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..3ec979003d836b3e24ceb29be619773bd9dbe790 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:867155a8f55351d5823c514b951994bb7b824162679aa765638eabb37ec86825 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a864fb90140ff24b5943f1c50a4664c58d356b59 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:254d1591e0d98b5784b977168b5df6e29521853c1a28bfd64b47bc392eb15532 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b7bdc8ae6d75c466dacfdfec9de371a70a50a01c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:acd7fc96a32453532d90cb8092d918262bc907f51f58d386c55ccf2349e1391b +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..f615af238cb04b2e5009c9cc7d843d3ad8a9bc83 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc204c94dcb0bd24b0a6d87a17fa53b76e99b19142b0f30ac2bc8aa4b0a984a1 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..df9770a1c9c8ce266628fc5a3d0b8c72b14a3f27 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:38d7079993650918dda6fc4674d2b66474d6664f031be972325357016eacdcc4 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..668d25a7331303ae1a801de8e2ae2e312ec1d1f0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6da2680b988d0e3ee10a6cd5692889bf6b2d7c2376f2f957fe2d46f9f30ee9b1 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..4de07907c0e7f050f8cad300977d8eaa1ca2e3ea --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63c5d402875024a6d430f8bb74338992ab9609405356ba3dc56711b637839dea +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..bb02235c880cd963d39e08ffd9e371c7fa64eb05 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/3_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.6424291133880615, + "learning_rate": 2e-05, + "loss": 0.8003, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.4942960739135742, + "learning_rate": 2e-05, + "loss": 0.3517, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.3215599060058594, + "learning_rate": 2e-05, + "loss": 0.2603, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.5830700397491455, + "learning_rate": 2e-05, + "loss": 0.6226, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.190831184387207, + "learning_rate": 2e-05, + "loss": 0.5444, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 4.842282295227051, + "learning_rate": 2e-05, + "loss": 0.7789, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.986421585083008, + "learning_rate": 2e-05, + "loss": 0.729, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.7310032844543457, + "learning_rate": 2e-05, + "loss": 0.5817, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.320521354675293, + "learning_rate": 2e-05, + "loss": 0.1708, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.1041675806045532, + "learning_rate": 2e-05, + "loss": 0.3381, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.876417398452759, + "learning_rate": 2e-05, + "loss": 0.2736, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.9649337530136108, + "learning_rate": 2e-05, + "loss": 0.4907, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.41735923290252686, + "learning_rate": 2e-05, + "loss": 0.2076, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 4.150988578796387, + "learning_rate": 2e-05, + "loss": 0.2896, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.223928213119507, + "learning_rate": 2e-05, + "loss": 0.4176, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.9266397953033447, + "learning_rate": 2e-05, + "loss": 0.3883, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.8420566320419312, + "learning_rate": 2e-05, + "loss": 0.6685, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.940491199493408, + "learning_rate": 2e-05, + "loss": 0.7417, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.357793927192688, + "learning_rate": 2e-05, + "loss": 0.1812, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.561838388442993, + "learning_rate": 2e-05, + "loss": 0.5488, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.2938289642333984, + "learning_rate": 2e-05, + "loss": 0.5713, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 6.213510036468506, + "learning_rate": 2e-05, + "loss": 0.6725, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.2259900569915771, + "learning_rate": 2e-05, + "loss": 0.2083, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.086206078529358, + "learning_rate": 2e-05, + "loss": 0.0714, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 3.5994632244110107, + "learning_rate": 2e-05, + "loss": 0.5986, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.135278224945068, + "learning_rate": 2e-05, + "loss": 0.6855, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.11214569211006165, + "learning_rate": 2e-05, + "loss": 0.2311, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.6004652976989746, + "learning_rate": 2e-05, + "loss": 0.271, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.4751770496368408, + "learning_rate": 2e-05, + "loss": 0.3031, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 3.720674514770508, + "learning_rate": 2e-05, + "loss": 0.5303, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.678459882736206, + "learning_rate": 2e-05, + "loss": 0.3599, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.482301950454712, + "learning_rate": 2e-05, + "loss": 0.2799, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.522612571716309, + "learning_rate": 2e-05, + "loss": 0.5165, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.6119441986083984, + "learning_rate": 2e-05, + "loss": 0.8158, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.7973575592041016, + "learning_rate": 2e-05, + "loss": 0.3723, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.715825080871582, + "learning_rate": 2e-05, + "loss": 0.6187, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.023906230926514, + "learning_rate": 2e-05, + "loss": 0.6829, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 4.166566371917725, + "learning_rate": 2e-05, + "loss": 0.2726, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.897386908531189, + "learning_rate": 2e-05, + "loss": 0.3291, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.9080281257629395, + "learning_rate": 2e-05, + "loss": 0.175, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.0254995822906494, + "learning_rate": 2e-05, + "loss": 0.1202, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.7732439041137695, + "learning_rate": 2e-05, + "loss": 0.205, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.966379165649414, + "learning_rate": 2e-05, + "loss": 0.5767, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0940295457839966, + "learning_rate": 2e-05, + "loss": 0.2869, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.4017220735549927, + "learning_rate": 2e-05, + "loss": 0.5045, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.67954683303833, + "learning_rate": 2e-05, + "loss": 0.3082, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.937507390975952, + "learning_rate": 2e-05, + "loss": 0.2375, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.4491846561431885, + "learning_rate": 2e-05, + "loss": 0.1879, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.3485398292541504, + "learning_rate": 2e-05, + "loss": 0.3423, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.9019141793251038, + "learning_rate": 2e-05, + "loss": 0.216, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5408561442062336.0, + "train_loss": 0.41872793197631836, + "train_runtime": 94.709, + "train_samples_per_second": 4.223, + "train_steps_per_second": 1.056 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5408561442062336.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..a8fbd33a6ab9037cd60e8d4f73097fb501e91b1c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eda6f96018f29050fd29080003b88d791f0173e2b0ab37c7db8cb6840d445670 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..857d9e099506e771f877fb7c6aa5a15889fd7279 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:741ff39960c145b788a1c0d3fc340ed7c6ea5eb162ea60901c07445f5d48132f +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..035ea839ae9f1d48b1592efa401249b9030b2be3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8f39923c03fdbb6eb7bde6e8fa15a07699fe8584cd79425fe1c71ee11d16be67 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0e98bfca81a6d0a0b75c79dbb1e83568e102291d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:da7eddaaf03a9cd7aa497e6beac519a6a76f9809bda69ff4d6e6c107291ddc1d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..21351f880468047f039c481e234dee0516d8549f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76c706fb4a3c9cf1ee81e64926849beb072205e9591c28d23f1e159205448186 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..2641da90b69c056ffc8248d07827842a32f81470 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0c39223ad99ef711c76760f5e47683820e93e80059e9ee6f992771305a81dfc3 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..55f30a9b17cd26af414489dc963b4711f08abb0a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2ecef05c19f92ddb73857abb30cc3ba519429861310697acabdde9d6b837be6b +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..12d397c0fd6e9a5e09ab9974de457a3ed95a9851 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:108553d91e7cd99f780cd7cc3ddf5707afadc3ea75cd25f11587a96229f4a970 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..927834bb0165d912a464a37891c3055b695c44b3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/4_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.0712028741836548, + "learning_rate": 2e-05, + "loss": 0.2024, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.7633742094039917, + "learning_rate": 2e-05, + "loss": 0.3808, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.2362414598464966, + "learning_rate": 2e-05, + "loss": 0.3528, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.129031181335449, + "learning_rate": 2e-05, + "loss": 0.3771, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.2611881494522095, + "learning_rate": 2e-05, + "loss": 0.2181, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.5939332246780396, + "learning_rate": 2e-05, + "loss": 0.1251, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.0693352222442627, + "learning_rate": 2e-05, + "loss": 0.3226, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.4935710430145264, + "learning_rate": 2e-05, + "loss": 0.634, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.7907732725143433, + "learning_rate": 2e-05, + "loss": 0.1703, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.523618459701538, + "learning_rate": 2e-05, + "loss": 0.2063, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.6210978031158447, + "learning_rate": 2e-05, + "loss": 0.9954, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.1934258937835693, + "learning_rate": 2e-05, + "loss": 0.4846, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.6618518829345703, + "learning_rate": 2e-05, + "loss": 0.3061, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.9813880920410156, + "learning_rate": 2e-05, + "loss": 0.473, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.8067651987075806, + "learning_rate": 2e-05, + "loss": 0.3786, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.6329550743103027, + "learning_rate": 2e-05, + "loss": 0.3718, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.9381539225578308, + "learning_rate": 2e-05, + "loss": 0.2056, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.220282793045044, + "learning_rate": 2e-05, + "loss": 0.4678, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.7268145084381104, + "learning_rate": 2e-05, + "loss": 0.2319, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 5.491344928741455, + "learning_rate": 2e-05, + "loss": 0.6719, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.5370978116989136, + "learning_rate": 2e-05, + "loss": 0.4108, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.8653494119644165, + "learning_rate": 2e-05, + "loss": 0.2812, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.548546314239502, + "learning_rate": 2e-05, + "loss": 0.3483, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.9249151945114136, + "learning_rate": 2e-05, + "loss": 0.3425, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.8268990516662598, + "learning_rate": 2e-05, + "loss": 0.2231, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.1500799655914307, + "learning_rate": 2e-05, + "loss": 0.7527, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.9101142883300781, + "learning_rate": 2e-05, + "loss": 0.3015, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.1729931831359863, + "learning_rate": 2e-05, + "loss": 0.1248, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.1908469200134277, + "learning_rate": 2e-05, + "loss": 0.3796, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.8858509063720703, + "learning_rate": 2e-05, + "loss": 0.4214, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.360217571258545, + "learning_rate": 2e-05, + "loss": 0.2396, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.6162331104278564, + "learning_rate": 2e-05, + "loss": 0.4739, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.6042280197143555, + "learning_rate": 2e-05, + "loss": 0.3611, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.9366803169250488, + "learning_rate": 2e-05, + "loss": 0.3577, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 4.1854023933410645, + "learning_rate": 2e-05, + "loss": 0.3721, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.8654505014419556, + "learning_rate": 2e-05, + "loss": 0.234, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.9890596270561218, + "learning_rate": 2e-05, + "loss": 0.301, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.8976471424102783, + "learning_rate": 2e-05, + "loss": 0.5818, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.8554496765136719, + "learning_rate": 2e-05, + "loss": 0.3279, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.7270612716674805, + "learning_rate": 2e-05, + "loss": 0.29, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.606137990951538, + "learning_rate": 2e-05, + "loss": 0.2349, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.711355447769165, + "learning_rate": 2e-05, + "loss": 0.3334, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.4700284004211426, + "learning_rate": 2e-05, + "loss": 0.1434, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.007743000984192, + "learning_rate": 2e-05, + "loss": 0.2135, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.4801454544067383, + "learning_rate": 2e-05, + "loss": 0.3672, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.264678478240967, + "learning_rate": 2e-05, + "loss": 0.2958, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.6900923252105713, + "learning_rate": 2e-05, + "loss": 0.4392, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.6152141094207764, + "learning_rate": 2e-05, + "loss": 0.3688, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.2249460220336914, + "learning_rate": 2e-05, + "loss": 0.0889, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 5.808964729309082, + "learning_rate": 2e-05, + "loss": 0.5781, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 6047770452426752.0, + "train_loss": 0.35529296875, + "train_runtime": 106.0841, + "train_samples_per_second": 3.771, + "train_steps_per_second": 0.943 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 6047770452426752.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5a189d0c9efb849f23459167c477cd53dbe12967 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:381d1a7159808aac7a5397229beb6b982fe1062326cb8697589f25386add4520 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..7fe5b04f892f0b93dbff4057ddde49ed0f7bb6c0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:77839a02fb95867276c0ce757ba6b970a43fa87d3c60b48946471c926ec1e9a7 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..3903ab091d314d14dfafb6b3874dcccbd5e72273 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bb94246fa0f9a3b1db6bb0c8809be000870190bba87f5470ea38249e0bbd74cb +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0b04ac97312b4d554c62fa78867b43a0cbcb4681 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d2929dd8562c5782dc82e2c48fef08abe0e1baf624bef62c0527852b519d56ed +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..1ea3984ea00fcdbd8a59269b6c692778dae2cce4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5eefa00218297b6b62224c5b9f55c1f1d2f44770a5988fef080b83587f064cd7 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..af455f1a7d071458da0336d0020f290c9eba1af1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebdef2109cbac393a4fe98cd5b8e6ae6d1ea42230657284f9a5a9c39f1e90ee1 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..83bd14511f8e9fd47c1d8b1d956b4ed1e6f4e170 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2722581341a8a32ed063fc947ed90c106ad58bf100013513b1664b243e72ec36 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..0fe633292e793fe955110cf69b3a0fbbeaf00127 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e4b9c5c22f8050750f0dbfcf4ec78764ab2d5e73941946df0221ef16a533393a +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..b3cefac77d769e3ea6505adab65154035666c086 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/5_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.8366425037384033, + "learning_rate": 2e-05, + "loss": 0.1694, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.683704376220703, + "learning_rate": 2e-05, + "loss": 0.325, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.6198794841766357, + "learning_rate": 2e-05, + "loss": 0.101, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.1528291255235672, + "learning_rate": 2e-05, + "loss": 0.1164, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.0297576189041138, + "learning_rate": 2e-05, + "loss": 0.1264, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.06911511719226837, + "learning_rate": 2e-05, + "loss": 0.0277, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.013462782837450504, + "learning_rate": 2e-05, + "loss": 0.0327, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.08065499365329742, + "learning_rate": 2e-05, + "loss": 0.0201, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.796327590942383, + "learning_rate": 2e-05, + "loss": 0.0889, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.047738552093506, + "learning_rate": 2e-05, + "loss": 0.1298, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.9473061561584473, + "learning_rate": 2e-05, + "loss": 0.0859, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.875561237335205, + "learning_rate": 2e-05, + "loss": 0.1445, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.6053429245948792, + "learning_rate": 2e-05, + "loss": 0.1855, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.5163557529449463, + "learning_rate": 2e-05, + "loss": 0.6293, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.14073553681373596, + "learning_rate": 2e-05, + "loss": 0.0395, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.8237165212631226, + "learning_rate": 2e-05, + "loss": 0.0398, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.074294328689575, + "learning_rate": 2e-05, + "loss": 0.1005, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3093288838863373, + "learning_rate": 2e-05, + "loss": 0.0377, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.188652515411377, + "learning_rate": 2e-05, + "loss": 0.0402, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.8078205585479736, + "learning_rate": 2e-05, + "loss": 0.135, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.7815762162208557, + "learning_rate": 2e-05, + "loss": 0.2799, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.437537670135498, + "learning_rate": 2e-05, + "loss": 0.207, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.20486502349376678, + "learning_rate": 2e-05, + "loss": 0.1832, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.12568062543869019, + "learning_rate": 2e-05, + "loss": 1.0219, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.753807067871094, + "learning_rate": 2e-05, + "loss": 0.1909, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.0715436935424805, + "learning_rate": 2e-05, + "loss": 0.1295, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.25005874037742615, + "learning_rate": 2e-05, + "loss": 0.0248, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.104605197906494, + "learning_rate": 2e-05, + "loss": 0.0997, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.2361832708120346, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.7761740684509277, + "learning_rate": 2e-05, + "loss": 0.04, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.37839582562446594, + "learning_rate": 2e-05, + "loss": 0.0365, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.28816556930542, + "learning_rate": 2e-05, + "loss": 0.5832, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.7323498725891113, + "learning_rate": 2e-05, + "loss": 0.0643, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.29171717166900635, + "learning_rate": 2e-05, + "loss": 0.0324, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.7399406433105469, + "learning_rate": 2e-05, + "loss": 0.2023, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.6951918601989746, + "learning_rate": 2e-05, + "loss": 0.2721, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.1435706466436386, + "learning_rate": 2e-05, + "loss": 0.2458, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.12317098677158356, + "learning_rate": 2e-05, + "loss": 0.0072, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.13525328040122986, + "learning_rate": 2e-05, + "loss": 0.5352, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.08037125319242477, + "learning_rate": 2e-05, + "loss": 0.0654, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.6846767663955688, + "learning_rate": 2e-05, + "loss": 0.0914, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.09761497378349304, + "learning_rate": 2e-05, + "loss": 0.0776, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.2910793423652649, + "learning_rate": 2e-05, + "loss": 0.0445, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.151703119277954, + "learning_rate": 2e-05, + "loss": 0.0491, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.4325088262557983, + "learning_rate": 2e-05, + "loss": 0.0371, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 5.531907081604004, + "learning_rate": 2e-05, + "loss": 0.2109, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.5840221047401428, + "learning_rate": 2e-05, + "loss": 0.0254, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.7832276821136475, + "learning_rate": 2e-05, + "loss": 0.2279, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.017747733741998672, + "learning_rate": 2e-05, + "loss": 0.2607, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 6.8001790046691895, + "learning_rate": 2e-05, + "loss": 1.3499, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289199494234112.0, + "train_loss": 0.18368536472320557, + "train_runtime": 102.6383, + "train_samples_per_second": 3.897, + "train_steps_per_second": 0.974 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289199494234112.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f87533791541846421600dd2912a3f948df840bc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f61b7fd261b96e56e815f73253b415eb72133fc83af9bca452c7082e4ec23e18 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..0e585ed17487866054542108167289c2ddb1e3f2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:05919e2892009b8d064a5f7a031240f7486b5d71e489ad61894a8c20917f4ac9 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..019b13fd9ba5a32c84d76e69e3bcd3ecd63e4dab --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a2034c200b162878350745fdc17ef6403830de0ff4fbb1e28e56b86fc5b676c0 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d955a4e03608efa479b283c26fbd6982f8873785 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c3fb0a80bd64f038b54bcd685e7f23218b802d1542239cb1f1a23678ecb513a +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..532854046422142a63e7c7c9f1d6cd926ab7611a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8a2ec9db5a8ba0f4d9f547f8a9bafe663fa816a2a5afb7c5bc2e12d4b7d44883 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d83602cf59123f312527ce28e152987e71b36f0a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56ec50e74e53bd40c5dc7b6fe773f512122f27d6a92b957d835a84fb5235ea82 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..9fb5dcb3082f21b77ba9f961a6807d882e0d0db7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3c15b8199979557b0d9fe7ba7605d44f3d0a580dad9bc63f2904bfcfb9fdbb5 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f1f918c55afe9d2c1ea864a50a1c68c650e78bbe --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4802c741252ceec25402005aa95364e2359ff739221ef0575e9cef66e4cbb1c0 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3cbdce4780b905589ae59102c536ff50480334f0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/6_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.312142848968506, + "learning_rate": 2e-05, + "loss": 0.3931, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.311955451965332, + "learning_rate": 2e-05, + "loss": 0.404, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.0048556327819824, + "learning_rate": 2e-05, + "loss": 0.4431, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.41806960105896, + "learning_rate": 2e-05, + "loss": 0.484, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.989250898361206, + "learning_rate": 2e-05, + "loss": 0.5425, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.48413586616516113, + "learning_rate": 2e-05, + "loss": 0.0394, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.8546371459960938, + "learning_rate": 2e-05, + "loss": 0.3205, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.226712226867676, + "learning_rate": 2e-05, + "loss": 0.5876, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.3076446056365967, + "learning_rate": 2e-05, + "loss": 0.7212, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.171445369720459, + "learning_rate": 2e-05, + "loss": 0.396, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.054872512817383, + "learning_rate": 2e-05, + "loss": 0.355, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.5660237073898315, + "learning_rate": 2e-05, + "loss": 0.2764, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.1519250869750977, + "learning_rate": 2e-05, + "loss": 0.3722, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.644512414932251, + "learning_rate": 2e-05, + "loss": 0.2749, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 4.26398229598999, + "learning_rate": 2e-05, + "loss": 0.8392, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.6505954265594482, + "learning_rate": 2e-05, + "loss": 0.3501, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.2055702209472656, + "learning_rate": 2e-05, + "loss": 0.4182, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.8087506294250488, + "learning_rate": 2e-05, + "loss": 0.5742, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.7071707248687744, + "learning_rate": 2e-05, + "loss": 0.4231, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.1986645460128784, + "learning_rate": 2e-05, + "loss": 0.4733, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.8725863695144653, + "learning_rate": 2e-05, + "loss": 0.8157, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.042865514755249, + "learning_rate": 2e-05, + "loss": 0.2329, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.469449281692505, + "learning_rate": 2e-05, + "loss": 0.3622, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.3576793670654297, + "learning_rate": 2e-05, + "loss": 0.321, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.1172313690185547, + "learning_rate": 2e-05, + "loss": 0.4089, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.1724525690078735, + "learning_rate": 2e-05, + "loss": 0.4028, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.878742218017578, + "learning_rate": 2e-05, + "loss": 0.3835, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.9443230628967285, + "learning_rate": 2e-05, + "loss": 0.8428, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.6403937339782715, + "learning_rate": 2e-05, + "loss": 0.4762, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.2627387046813965, + "learning_rate": 2e-05, + "loss": 0.5663, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.531160354614258, + "learning_rate": 2e-05, + "loss": 0.3591, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.6137160062789917, + "learning_rate": 2e-05, + "loss": 0.4417, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.0742682218551636, + "learning_rate": 2e-05, + "loss": 0.5122, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.019188404083252, + "learning_rate": 2e-05, + "loss": 0.411, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.9247498512268066, + "learning_rate": 2e-05, + "loss": 0.668, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.0114731788635254, + "learning_rate": 2e-05, + "loss": 0.3409, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 6.340407371520996, + "learning_rate": 2e-05, + "loss": 0.6909, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.640566349029541, + "learning_rate": 2e-05, + "loss": 0.5688, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.0275142192840576, + "learning_rate": 2e-05, + "loss": 0.2771, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.004143238067627, + "learning_rate": 2e-05, + "loss": 0.2328, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.631279468536377, + "learning_rate": 2e-05, + "loss": 0.4479, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.6522293090820312, + "learning_rate": 2e-05, + "loss": 0.6257, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.8231219053268433, + "learning_rate": 2e-05, + "loss": 0.2626, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.5955970287322998, + "learning_rate": 2e-05, + "loss": 0.4355, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.02359676361084, + "learning_rate": 2e-05, + "loss": 0.3196, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.8789767026901245, + "learning_rate": 2e-05, + "loss": 0.9409, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.518844485282898, + "learning_rate": 2e-05, + "loss": 0.3152, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.3316898345947266, + "learning_rate": 2e-05, + "loss": 0.3276, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.42132306098938, + "learning_rate": 2e-05, + "loss": 0.4854, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.543337821960449, + "learning_rate": 2e-05, + "loss": 0.9688, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 1.0467000203608064e+16, + "train_loss": 0.462637939453125, + "train_runtime": 163.4445, + "train_samples_per_second": 2.447, + "train_steps_per_second": 0.612 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 1.0467000203608064e+16, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..af29558e9a80a72b4870e4bd3b585a71427ca65b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c1c2cb07e6d64001668bdf88f35f15e68020fd602c3737b0b1820d8ff412daaf +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..47b44f0e027be3bb92d587cd0399ff98df552d8e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:16c4e9dc168b653510fdc0fff01a99ac63144b85c65d89bd0ba9a5829939445d +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..cd4c701aff6d83a7c78a9792e411f16e0576f924 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ed9d38532052c148c25fa7e5150a8817c86b704c6e79b1041c633d19000ef96a +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d2adf44e2989786ce2fe0be0fc3698cca87a4886 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9de2c05ba8439c56835d27141a18ef02771e47c079b56568318cb45efdd51228 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..4e51b60701e001e2755b8215ec3e2e082d643f73 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bded2e0cc4bff78430f5877008e8480f2d8b4065059bed7a8799f170639172e4 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..e73d8994344666435924e7ed85ba7bc233bdb2a1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ddf8bbde2352770bdf520e6d0c00259e197e2134c01d55919e7bfee5f1aa6c98 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..48958e8fc8b2a6b1d419792a49aa759cb71ddca2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5a1ab5686eea82b39cec91410c5c09fa17024b6c8192cba52805d57391aff2f7 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..5d957d749da1fe3a707a4b32d25ec4ec06aec48f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a79120db443b0607867bdbadf4936af8c395fdf736549ab884494620982f77b9 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..5bcbf698f9f85f3605ec2843f1d0d13f79be3d11 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/7_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.1755051612854004, + "learning_rate": 2e-05, + "loss": 0.1122, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.634487628936768, + "learning_rate": 2e-05, + "loss": 0.4212, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.681210517883301, + "learning_rate": 2e-05, + "loss": 0.1804, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.8976295590400696, + "learning_rate": 2e-05, + "loss": 0.2076, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.3863887786865234, + "learning_rate": 2e-05, + "loss": 0.7006, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.11061318963766098, + "learning_rate": 2e-05, + "loss": 0.0418, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.3342927396297455, + "learning_rate": 2e-05, + "loss": 0.0943, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.73251473903656, + "learning_rate": 2e-05, + "loss": 0.43, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.4592249393463135, + "learning_rate": 2e-05, + "loss": 0.1916, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.077377796173096, + "learning_rate": 2e-05, + "loss": 0.3045, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.156890869140625, + "learning_rate": 2e-05, + "loss": 0.2495, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.3694241046905518, + "learning_rate": 2e-05, + "loss": 0.2715, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.6508183479309082, + "learning_rate": 2e-05, + "loss": 0.0346, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.6611247062683105, + "learning_rate": 2e-05, + "loss": 0.3232, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.4086673855781555, + "learning_rate": 2e-05, + "loss": 0.3432, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.9354144334793091, + "learning_rate": 2e-05, + "loss": 0.2638, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8759964108467102, + "learning_rate": 2e-05, + "loss": 0.0412, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.2730793952941895, + "learning_rate": 2e-05, + "loss": 0.4463, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.8122385740280151, + "learning_rate": 2e-05, + "loss": 0.0623, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.628685474395752, + "learning_rate": 2e-05, + "loss": 0.6793, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.5241402983665466, + "learning_rate": 2e-05, + "loss": 0.4718, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.73521089553833, + "learning_rate": 2e-05, + "loss": 0.2319, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.6701934337615967, + "learning_rate": 2e-05, + "loss": 0.2878, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.2300349473953247, + "learning_rate": 2e-05, + "loss": 0.0097, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.3767426609992981, + "learning_rate": 2e-05, + "loss": 0.4195, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.3898110389709473, + "learning_rate": 2e-05, + "loss": 0.5038, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.03246752917766571, + "learning_rate": 2e-05, + "loss": 0.3116, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4891020953655243, + "learning_rate": 2e-05, + "loss": 0.0296, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.1901887059211731, + "learning_rate": 2e-05, + "loss": 0.1013, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.532278060913086, + "learning_rate": 2e-05, + "loss": 1.2327, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.5174512267112732, + "learning_rate": 2e-05, + "loss": 0.0569, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.9170253276824951, + "learning_rate": 2e-05, + "loss": 0.4746, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.618475317955017, + "learning_rate": 2e-05, + "loss": 0.1471, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7863269448280334, + "learning_rate": 2e-05, + "loss": 0.1195, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.67844557762146, + "learning_rate": 2e-05, + "loss": 0.3609, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.7931462526321411, + "learning_rate": 2e-05, + "loss": 0.3013, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.6811180114746094, + "learning_rate": 2e-05, + "loss": 0.4751, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.5165439248085022, + "learning_rate": 2e-05, + "loss": 0.3907, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.752575397491455, + "learning_rate": 2e-05, + "loss": 0.3736, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.9842543601989746, + "learning_rate": 2e-05, + "loss": 0.8044, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.477774739265442, + "learning_rate": 2e-05, + "loss": 0.1654, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.6356291770935059, + "learning_rate": 2e-05, + "loss": 0.1201, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.1548612117767334, + "learning_rate": 2e-05, + "loss": 0.4333, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.139477491378784, + "learning_rate": 2e-05, + "loss": 0.3884, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.8254485726356506, + "learning_rate": 2e-05, + "loss": 0.2675, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.9783837795257568, + "learning_rate": 2e-05, + "loss": 0.1552, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.0617423057556152, + "learning_rate": 2e-05, + "loss": 0.1082, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.2474711686372757, + "learning_rate": 2e-05, + "loss": 0.0728, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.4221913814544678, + "learning_rate": 2e-05, + "loss": 0.1823, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.060051441192627, + "learning_rate": 2e-05, + "loss": 0.263, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5500391349288960.0, + "train_loss": 0.29318286895751955, + "train_runtime": 116.6192, + "train_samples_per_second": 3.43, + "train_steps_per_second": 0.857 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5500391349288960.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..38cdc1e4facd2d83cfb93d1ff8c2104fc5120122 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c8144a9dc3817a5824d4b897484dca4359896deed812f6d3d89f8d13f7b0b163 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..45da15a10ebbf5fc223a126e468ea56a6dd514d6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4c8d43a4d62f653d0219d2e619491369a4cad55294a9f81f0f3dc704a6afab7c +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..3ee3fb312ef68d157b108480e23aebfc864ba22a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5bb31333c0359ee7930005d29b8bdc967b003ab3a8228e68a7b13077ef5cce6a +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e8ea7316f900e3d3d4726c414a50d5109b02dedb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0024a606ec57feb4cc72a5863fef2921b044d36396ceccdaa7b8fe5e0fc80b73 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..034d7d66443439ae17dbf0a80cc3ed92ac4d4222 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f295bfd3cf183c0e2adf4a79ed6827a336279d7f36fb29b33c6bc30daa4acd89 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..360d07a9854b050e0e80cd910b1d0b5b9cf3998e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a964379317741084877edaf52b6ff5e3d70780045377d7af91ca1a6cdb208f73 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..a9ea33cdee4caab41a92a924cc82d7c75548659f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6ffa9ae7772dfcd801a6886b0d2ac1d26f477b81a8523e16d86d6ab4c5b24704 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..c436bf70847951b9dcc353df26ea10520349f0f9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:366ac33636b35d589522eca87e03edde7bc89c736a46f550a063b5efb379eae4 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..02d8e2a090e54f21f598323600ec77dbc1e8567e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/8_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.027203630656003952, + "learning_rate": 2e-05, + "loss": 0.0246, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.06620825082063675, + "learning_rate": 2e-05, + "loss": 0.0739, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.1927112489938736, + "learning_rate": 2e-05, + "loss": 0.0439, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.415907859802246, + "learning_rate": 2e-05, + "loss": 0.2692, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.03882341459393501, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.7003546953201294, + "learning_rate": 2e-05, + "loss": 0.0908, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.06669411808252335, + "learning_rate": 2e-05, + "loss": 0.1466, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.604789733886719, + "learning_rate": 2e-05, + "loss": 0.262, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.4016295671463013, + "learning_rate": 2e-05, + "loss": 0.1111, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.369394779205322, + "learning_rate": 2e-05, + "loss": 2.2642, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.52653169631958, + "learning_rate": 2e-05, + "loss": 0.5513, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.6306421756744385, + "learning_rate": 2e-05, + "loss": 0.4034, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.9526010751724243, + "learning_rate": 2e-05, + "loss": 0.1058, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.1774262189865112, + "learning_rate": 2e-05, + "loss": 0.0549, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.3812141418457031, + "learning_rate": 2e-05, + "loss": 0.0495, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.06561172753572464, + "learning_rate": 2e-05, + "loss": 0.0096, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.288865566253662, + "learning_rate": 2e-05, + "loss": 0.6997, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.8393924236297607, + "learning_rate": 2e-05, + "loss": 0.3152, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.39931583404541, + "learning_rate": 2e-05, + "loss": 0.5663, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.6632494926452637, + "learning_rate": 2e-05, + "loss": 0.6159, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 5.843943119049072, + "learning_rate": 2e-05, + "loss": 0.5173, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.0021515991538763046, + "learning_rate": 2e-05, + "loss": 0.0606, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.4243650436401367, + "learning_rate": 2e-05, + "loss": 0.3904, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.8705551624298096, + "learning_rate": 2e-05, + "loss": 0.3873, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.3426576852798462, + "learning_rate": 2e-05, + "loss": 0.1898, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.025865498930215836, + "learning_rate": 2e-05, + "loss": 0.2212, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.30158159136772156, + "learning_rate": 2e-05, + "loss": 0.2386, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.3702574372291565, + "learning_rate": 2e-05, + "loss": 0.2925, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.2034703493118286, + "learning_rate": 2e-05, + "loss": 0.0722, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09868835657835007, + "learning_rate": 2e-05, + "loss": 0.0096, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.4942784309387207, + "learning_rate": 2e-05, + "loss": 0.2023, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.221104145050049, + "learning_rate": 2e-05, + "loss": 0.5935, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.308194160461426, + "learning_rate": 2e-05, + "loss": 0.1152, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.8967278003692627, + "learning_rate": 2e-05, + "loss": 0.0327, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.5046753883361816, + "learning_rate": 2e-05, + "loss": 0.0243, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.801755428314209, + "learning_rate": 2e-05, + "loss": 0.1714, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.7808038592338562, + "learning_rate": 2e-05, + "loss": 0.0308, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.6014570593833923, + "learning_rate": 2e-05, + "loss": 0.0419, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.2914716303348541, + "learning_rate": 2e-05, + "loss": 0.1087, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.2409433126449585, + "learning_rate": 2e-05, + "loss": 0.0192, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 4.595032215118408, + "learning_rate": 2e-05, + "loss": 0.4703, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.6504670977592468, + "learning_rate": 2e-05, + "loss": 0.0255, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.18435387313365936, + "learning_rate": 2e-05, + "loss": 0.0092, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.8855642080307007, + "learning_rate": 2e-05, + "loss": 0.0449, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.767470121383667, + "learning_rate": 2e-05, + "loss": 0.153, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.6871795654296875, + "learning_rate": 2e-05, + "loss": 0.4686, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.22215917706489563, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.82851243019104, + "learning_rate": 2e-05, + "loss": 0.1452, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.2504313886165619, + "learning_rate": 2e-05, + "loss": 0.0095, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.15292932093143463, + "learning_rate": 2e-05, + "loss": 0.0085, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5327278363901952.0, + "train_loss": 0.2344178366661072, + "train_runtime": 110.7172, + "train_samples_per_second": 3.613, + "train_steps_per_second": 0.903 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5327278363901952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..a13935d8689680e3745505c5ae130465e34691f7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70503cf4624ab33a09793064346b11cb0685ba6e8a3198364884bde67e6feb82 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..57724a989994ec458526e976752df1948e6aa063 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3e6fb53faf99bff6290cec74ea3de1575a5ff383fc4d1d5a46a0449c1c1fa700 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..99bfb015273427ea05ba7931dc07fe633a0523b8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4595a41193f4aa0f8fdc17779e77c356302b4b8b152544e1a5c87fc9f277f7cd +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..7a67a09af028bd73e38f14cba814deaa1a3e397a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5aa917f89c23fa4305b68918f2098041706c5bf1624076d4f75c97cbd10d3ae9 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..3599a2c43bf8bc6c9d2441c7c065cad7206bfa7f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:02bc9784a7182d089937f67449ac1d90cab85cc6edb527246b4cbf5195deca7d +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d8880d76fe7edeb0e1300c617528a095491869ea --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8135f9a70376a4204dfaf1c9a7aecd9837b03a93ecc0bb298ebe003d0f71c112 +size 791578182 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..1da326bcc49e1c5e9e37119e3135303f0db0f5ce --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0cb4000c723f3837156a0af3597744183d8ce084e72bd1b8afa71fa8ca401f41 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..42e0421e214f4d3b4591cab703631f2f14d956ce --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:60a8d346c3b8eb2e2ce447d8c1c1f89b9e3a010968d864779662bb855db60142 +size 791576546 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..1432f7025b2a197b8dab5c728bf3df97f852aedf --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/9_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.2871222198009491, + "learning_rate": 2e-05, + "loss": 0.3477, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.6915421485900879, + "learning_rate": 2e-05, + "loss": 0.6724, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.9630380272865295, + "learning_rate": 2e-05, + "loss": 0.0801, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.3849560022354126, + "learning_rate": 2e-05, + "loss": 0.1112, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.145646870136261, + "learning_rate": 2e-05, + "loss": 0.3151, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.6976999640464783, + "learning_rate": 2e-05, + "loss": 0.1491, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.08228840678930283, + "learning_rate": 2e-05, + "loss": 0.3653, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.594034194946289, + "learning_rate": 2e-05, + "loss": 0.5891, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.8721482157707214, + "learning_rate": 2e-05, + "loss": 0.1492, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.0610108375549316, + "learning_rate": 2e-05, + "loss": 0.2297, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.793205976486206, + "learning_rate": 2e-05, + "loss": 0.7183, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.994863510131836, + "learning_rate": 2e-05, + "loss": 0.4926, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.9002552032470703, + "learning_rate": 2e-05, + "loss": 0.4739, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.362318992614746, + "learning_rate": 2e-05, + "loss": 0.0794, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.2748732566833496, + "learning_rate": 2e-05, + "loss": 0.2208, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.14847247302532196, + "learning_rate": 2e-05, + "loss": 0.006, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.852799654006958, + "learning_rate": 2e-05, + "loss": 0.5546, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.2233632802963257, + "learning_rate": 2e-05, + "loss": 0.1693, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2944558560848236, + "learning_rate": 2e-05, + "loss": 0.0618, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.635985851287842, + "learning_rate": 2e-05, + "loss": 0.3375, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.2670395374298096, + "learning_rate": 2e-05, + "loss": 0.3085, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.843581676483154, + "learning_rate": 2e-05, + "loss": 0.6163, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.0938365459442139, + "learning_rate": 2e-05, + "loss": 0.0656, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 4.283212661743164, + "learning_rate": 2e-05, + "loss": 0.4414, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.9494190216064453, + "learning_rate": 2e-05, + "loss": 0.0802, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.9024108052253723, + "learning_rate": 2e-05, + "loss": 0.095, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.7638759016990662, + "learning_rate": 2e-05, + "loss": 0.4953, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.1802085041999817, + "learning_rate": 2e-05, + "loss": 0.0162, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.8723180294036865, + "learning_rate": 2e-05, + "loss": 0.352, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8824854493141174, + "learning_rate": 2e-05, + "loss": 0.0955, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.0696358680725098, + "learning_rate": 2e-05, + "loss": 0.1786, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.5412893295288086, + "learning_rate": 2e-05, + "loss": 0.2386, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.669694900512695, + "learning_rate": 2e-05, + "loss": 0.2339, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.0899804830551147, + "learning_rate": 2e-05, + "loss": 0.0859, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.12675884366035461, + "learning_rate": 2e-05, + "loss": 0.0153, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 5.6869587898254395, + "learning_rate": 2e-05, + "loss": 0.4697, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.9030308127403259, + "learning_rate": 2e-05, + "loss": 0.2303, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.810260534286499, + "learning_rate": 2e-05, + "loss": 0.3415, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.342296838760376, + "learning_rate": 2e-05, + "loss": 0.1485, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 4.620863914489746, + "learning_rate": 2e-05, + "loss": 0.4314, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.693507194519043, + "learning_rate": 2e-05, + "loss": 0.2811, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.982360363006592, + "learning_rate": 2e-05, + "loss": 0.2211, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 3.461449384689331, + "learning_rate": 2e-05, + "loss": 0.2182, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.037543773651123, + "learning_rate": 2e-05, + "loss": 0.3811, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.0902976989746094, + "learning_rate": 2e-05, + "loss": 0.2746, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.766568660736084, + "learning_rate": 2e-05, + "loss": 0.3017, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.131326198577881, + "learning_rate": 2e-05, + "loss": 0.5801, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.6314743757247925, + "learning_rate": 2e-05, + "loss": 0.054, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.250058889389038, + "learning_rate": 2e-05, + "loss": 0.2443, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.2442658245563507, + "learning_rate": 2e-05, + "loss": 0.0582, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291656215527424.0, + "train_loss": 0.2735392713546753, + "train_runtime": 102.9729, + "train_samples_per_second": 3.885, + "train_steps_per_second": 0.971 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291656215527424.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ccf7429312c8eb302a8722ac1a124d7ccc842ced --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f2883aab19e0a2a98dbc78aed7b4504784d25363622e42e73fcf98196c3c41fe +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..8b336b3ba62385278b1afd77bd2050f5d443bc92 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:350d68ff39489d43ee594198a679f1ede80ac0bf85f3a4fc1fdc03c719e03f1a +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_0/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d13fc1b9790fd3c7410f7c0458a52fb85b5b408a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:039ba4e834a7dd691028aad5bf1ffca0a7bf2c3446f454adc3398c59768a5dc3 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d631ca6adba2c3d4dbb86e26d0be1abb29966f42 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:48aa9e8253ea22f255b70c97c3d1479fd362f2e5d57d843df56950743da75eb7 +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_1/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..73077e495c61c0f5e8287ae75ef2f2178669e49f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b018c7c128c092ba2824b7608755e87c899d8554578ccc8f64e36a5fab6a2c93 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..20d13bbcfa9b1fb3841224a44aaa6c66f1dbcc9b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd38ef50476f62ee7f4b265710173990696e0966a040a1f33a40b3a1de53e3f0 +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_10/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1f2e42473311e2c605263fd4da1eea3fddddf043 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:52705fea7afc48d31d1484c41a66d846c68fe3c2742bc1b005032faf8d95cce8 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..fd025e9381bc45a4f1f34a18683f57f1125e258b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3b9e8a5f7fce607b86d34d651d9111426545b8d276523e1f37f43d6077bb1959 +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_11/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..15156eb5bc0029ac794be56277312103501a0af7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1d5f29a5846769a38c0f85027e269176dca80e5995fe27b45c97e62fa76ad44b +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..84b9ee0ab5c1486c3ae368d144a95ff523b79dde --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ed57bb8f5e093dd218cad9f7e6ff8e53b77c57a5f4fb903dd678fa21031f66af +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_12/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c75bbf81fc091d74170c540966be4f85d3a73305 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a21b3bc74a40cd8cc8434648c1999f13c9eb3c3f7cb1f2aee7d5b72dae58ce6c +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0e499e8a3b7f017e40d9421c6e4159fefb8a37a7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1b5e0212369d5b2a1af351bf203c43c7d9ecf1acce218b875aef17c19e59e423 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_13/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..74a150b5cf051c803498923a5a99ade8202f0c97 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:77187d098ceb76bfae8be0b6cd5d7f202c5ab8fcbf70c12ed1525a0a1fe26688 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b5bb874001e0a1d093ce2482bcfe1a1cd938bb89 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2d47bd9e4b16b259740153eec14731009dc2275d7f342ee69ec2d55c7e329928 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_14/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1290f10f0da54404e80f00c4aff2f7f6f5a343da --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3cf959b18ed0a6a3b483119e529c7dee67165ba306ec518ed6fce4d5bf7dc1d +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1c4ff13dd2b4d0592c1bdea0d1a463e89868d4e7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9b05e99ddbfdaff2e26bcfd4d0d9b8136121fca15cd142f025248c664d852a3b +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_15/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d358dd56aa93b4b92852a50f9295d6a049e6bd58 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6db0d7aeef959ab80a2b7ceec4481b666ef012515925cbce63381d9270e49842 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..48b4ce5edc1d9d516f9d195846281edae6adcad5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c51d775b7acb9b61d4053825c2f4bc4bf89e814a6ab0c72a53fd4cff7f4b536 +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_16/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2dfa860a375e352422c8b3305a48c8b89d14dfb1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5fda1faab75c6b08004afab0444364f4f5d49cbb8b8bf71848f3dd08530c9a9c +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..35eda829703bac2933bfcee7bf961ac4fd84fe3a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d4f18d1cd4893a3f0a0877ce46d162406771e01f9474e097743ea48628e560f6 +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_17/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..62ef1b7b27872c7bfc053abee680329fc87aac93 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7e48f8f46f99f6fa41d5142d1276e7c4a490d7f9510d80fe9d0931eef36ca279 +size 969111984 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..8fe614b20a7a2ba8da6c85ba4238a292bc501636 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03beda0ca64c1775a0e722a711a0c328b4d1ad5f30adcb7cbd8100e7afba2e9b +size 894034212 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_18/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..188ce715cffca2a18e1286ffb4404400513fb7f1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:78bfe45949c1a0e43989e0dff18ee9b65476ac402e75d1256916423c6b3ca4c1 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6e7e3de7b265e240752f81253f8c1d81ccc8aba2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d19d9c0ad8247939418837fee53767a03cc767e2711fe89d26a791e54a97e25 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_19/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6a571b5a25bba7f843a6e3f2e293574b06401ae3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f9abafa9ca742f1f8801dbc8ceb41f7fddd074dfd5d4dd5210cbb55bea78ffb9 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..244f9eb2173dfa6a8006bc3b7c5b9a4c8d9800a0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bc3093ddce0348f54e012f17370d6a1573f32d2d4cddf025d5cc565abbf1bac9 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_2/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7bcbaacbfb8a5c6a24ffd79c5a09efb90db13021 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:40b53e959e59171f4d7cd98e4dd92ce5d11d752af87be651d4ccdd3e07918993 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f991fbd609008a2a554a31cb2c6edc599c96c976 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c4f0cce6516faeaa66b57699e8e059fdf1783ec1c3adbb74e83009370d18b63f +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_20/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e7e54a21f4eab62aaa15e9ac68a7b71cec924c7a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4acafbbd4de4d5665f53b6f9b25b5107a10cc94c99c0236c517bde749adf97cc +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b572cc8e3ff337d53c88dbb9c220f922a8a7b01b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a2cb33967b90a29142e946cb2f81713e0ca213673b501a4c2e285bcf202f143a +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_3/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..af957e9dfdf6da2c5ce52e13e62212b1ffff04a3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3adac8b2d38e26877fe2fe0d21eab290c0a116a5296199e10a5f149b12389010 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ce8ec9e654918a754992d948828f9158b06b136d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:978e1f6f8c42e85622c36c99bbafafa837695c6b75518c1eb5da873fcc491998 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_4/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c09115b43fe412d102ef099c25b3430ebe33aab9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:610b2399c2a78e4be936950b9805e492201910889f80a707fabc202bef9e43eb +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e1bc205c70af6dd7a89d5760663f30ab3ab86f78 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0927a94ac959834ffe76f12a4c194f116abe3bedd0c4733943c29dc93c3e8ffc +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_5/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..5f8187060758aae17274f083d66dca1111f1baf4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8a3a88ede04e00b08c79926030ba88e616d4139d19c7ba47a5bb070cda8b51bb +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..90ba05fbaa274ebfaeeda644394957464f18295b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fde8f16874aa3c8c5bbfad1c105464f9c428099469e6f84cd472dd0337c6ab16 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_6/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..38b0f518e45dfb2ca083b7da5c16a07dbd4d0bc4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:834c5332904b47f27f38973f2389c9be3ba5a06b3258d653a80c190f4ca62f00 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a379293705296f8616dbf58ee5c91f129a720eee --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3fc90d80f39c268e406f5e30ecb2b28e474df3deff4199108ce3bcfdd769876 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_7/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..01879ca85d8705f7428fb064550c11c6df8258cb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c6421f928f0fd1f624f01045ed93fd2b9c9ee6c58dadcf60a46088f4de1c1ec4 +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c7dcd450e3fe2525cf4203ebe8f9c69df5bae3dd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:785cc877564810942a333f1cbf79372ea21387189d236e711157096014ccd5b7 +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_8/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..8f4314032b68463f984f6f97ba43e766dca72c0c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8a66d25f592fff62bb8c0c8532bfbaf81709ddc729b1df0cbbe741fb58f86a5b +size 2206339760 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..91cf4aecb4c7b103521bd1fac6be62542a1c52bf --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:782b2b11358d70638073b336a2f47e9579a36113dd51f48e16d25d5de829214c +size 1156402372 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/client_9/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round10_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round10_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..8480f3290fe02dcc8ad3f411dce4bb771607f136 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round10_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2a7a8d47edc80b205a0230dd0d312c16f3cc81aa0b01bd0d21718b6314735bb +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round11_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round11_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..120369764554a4d6cac58078d8512bf359607829 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round11_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf1dde35905b1ad125edf778aff1f9c3576ae0715c8c9876f665a63efcf18cd +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round12_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round12_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..64b454d945aa2a842c509c5cebe5ce0706d83a0c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round12_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb23ed1efc2080025ddd7ca376ff771b3f058afa24e048f5cb3cf99201aa1e03 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round13_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round13_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ad8d07954c41361e73a0c9c7d3b4b46d894b188 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round13_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8caca9dd439d57dbc5054f65d33ebbd7ac694940b2b9dbe5e8cdf68881137a2 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round14_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round14_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f05da1a0577050cea961f470b13842ab5f1d3f9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round14_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafca5706ac84f575d7648ee4786543561523e58cea8c01d2214b03d300bb988 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round15_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round15_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..351a4d2a284717aec90435f50067e1bb7002e9e4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round15_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:341a5a225c4c17a621d4f60caf4c34834000ea74e6f3c8b4f5aaada72ca1f547 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round16_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round16_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..08cfb19e58b68e297c605b87165d4acd241054e5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round16_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb378347913e2dd7be50fa6f9e3bd9e481f3897c2c9d31c597fe9164efc50bee +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round17_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round17_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..aeee08397c6123d33648f6703464379a4b7f85fe --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round17_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70c677ae195453774b8815440fee856ddb773c322681ac422ebc6e8abd5a58d7 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round18_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round18_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9935621fdc6da9e4abad379d97303b5ebf4f8f5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round18_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:393c3d5a15b40b2ff049c254f5b09f25d226b0971827b4e7460ca58f9ff7b76f +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round19_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round19_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..2766cf1a4da1390df01ea4146b4a3a7a06e1434e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round19_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dc77db94ac5b4746eebf2c850a9425da9a425a71108790d6b7f5f2654b614a8 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round1_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round1_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..34bfc5121e29fc587c3f574c17d8006189d3021a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round1_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47a94b88b812f21101345aae20a53cc10e0b346660c2cef669c77b201d06d508 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round20_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round20_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..4be93ba48f85d89c8cbd0e7426b79a10f51bc786 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round20_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20b9a2e67fdcb4051cebf52cba2ed0d47facb4dc69170f35a3483e3601bf25a4 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round2_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round2_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..b42d717305b30cd186c1c408b8719020c0658a5d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round2_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fd7458877767da0a36b80344cd6ec7eae8fbc8e11d07f8bfce3b4d5f5f239a9 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round3_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round3_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..01842d1316f21c3326c5735c6257b499da7ca458 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round3_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e537f056156082569ae0819abd06beb04b46581d0c63843bd68926e0ac2bff0e +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round4_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round4_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..3af9fe115499486166c3ef6d0a7bd27661ac08b4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round4_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a9649d6997f84508e5f566393e91e42e03a31aca55878bacd7ffd2e134bbfc +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round5_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round5_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..abd1e557c23a8769870b0beedaf7df93836a101e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round5_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7520bace4f501771d879621facd1b44f8db1ad7385a653e1535c7642465e67d +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round6_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round6_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..ebc846b7a5c587ced387cca85c4065e921c8cbe3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round6_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1951a910d1b7587459db2110e8ac43cbc906a34267372eb21d2244a900eaec46 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round7_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round7_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..5eb1370d3fa3177cb4fb710882f4b7eb706817e8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round7_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03f40a0df2437812b9e6de4e5d4e30014acb71344efa8132f19e1fc0f5925c82 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round8_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round8_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..385c599585a0973f240a027d28e985ca72d6b721 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round8_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3ef6d8247bc37b017895c9e5bbaf32f2db500f94f66b4ebade5a4a5ba7a271 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round9_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round9_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..219adb8927937a4597cd82efbf7ebb4c8a7519ad --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_iter100_round5_hetero/round9_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:277fa5f0607d993e80153171e4760f999c4ee6a65452ea62694d0b9b03818f22 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..c1838e4adfab41ccb9fe869958cc0d878be9ff78 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8024a5444a0c10afee93a14b518b7996fb9cd912ce2220d49bc509b28f1780f1 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..09564d42d67193c5dd5915fa5d6ab7763059b1b5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:77d5f7a895b357e68bbcc7fec64acab6f9a7f268f8ab99bf384a4593c3a87e3c +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a8663a0fd30309549396a7864fa5c22dfc206e3a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93c8b7400ad62ffb278f6f054b97d35c0524424d70516509d3072e0a6f34f727 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..984f2f348e3bb289779f0e99c355ab1b9d9ec629 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b1724b7644e4ca779445522020ae9cb4eb43dd3c15ef647943f859689c0d47b4 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..309673c61a53775704c938519464ee509d9facaa --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c4a975e9d6cd2c1c46897ee6c6217bb95b65e7ea5417941d816ee56571f4c1f +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..05a4731ef4d4a9394fb717ba2e87798660b2a534 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d340fd3ad76e02dca1db64a7d4dedd98a19e6cb4daca191d7fde7a95fce85837 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..80fc7e887109a1078d5776c765ae9f8a81d6f0e9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:faafadf85a023d03a9eefe7458b8540dd9133b5c29322914043ee4631d22ab0b +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d4b6240151a01cd6d3ccc2feff5f3bc018d6beb2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d044c336bba20c87dd81605f7516631a00e95fee50f15dab192af801c7bfe880 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c15f12b547c8739199739aa14f9fed9d6db1f9dd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.859551191329956, + "learning_rate": 2e-05, + "loss": 0.1643, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.8229276537895203, + "learning_rate": 2e-05, + "loss": 0.1669, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.32890793681144714, + "learning_rate": 2e-05, + "loss": 0.0306, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.16120509803295135, + "learning_rate": 2e-05, + "loss": 0.0373, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.945129632949829, + "learning_rate": 2e-05, + "loss": 0.3054, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.407341957092285, + "learning_rate": 2e-05, + "loss": 0.2358, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.8118242025375366, + "learning_rate": 2e-05, + "loss": 0.0906, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.010066270828247, + "learning_rate": 2e-05, + "loss": 0.4678, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 4.013240337371826, + "learning_rate": 2e-05, + "loss": 0.9585, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.2035651206970215, + "learning_rate": 2e-05, + "loss": 0.3077, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.0824318453669548, + "learning_rate": 2e-05, + "loss": 0.0109, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.708960771560669, + "learning_rate": 2e-05, + "loss": 0.394, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.22119958698749542, + "learning_rate": 2e-05, + "loss": 0.1896, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.2411755472421646, + "learning_rate": 2e-05, + "loss": 0.2326, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.2401578426361084, + "learning_rate": 2e-05, + "loss": 0.2777, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.06962911039590836, + "learning_rate": 2e-05, + "loss": 0.0201, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.2609751224517822, + "learning_rate": 2e-05, + "loss": 0.1671, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.072238445281982, + "learning_rate": 2e-05, + "loss": 0.3199, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.5981231331825256, + "learning_rate": 2e-05, + "loss": 0.2571, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.8929151296615601, + "learning_rate": 2e-05, + "loss": 0.095, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.17781028151512146, + "learning_rate": 2e-05, + "loss": 0.019, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.094531774520874, + "learning_rate": 2e-05, + "loss": 0.3066, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4083300232887268, + "learning_rate": 2e-05, + "loss": 0.0984, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.276772141456604, + "learning_rate": 2e-05, + "loss": 0.2035, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.7436290383338928, + "learning_rate": 2e-05, + "loss": 0.0764, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.1597166061401367, + "learning_rate": 2e-05, + "loss": 0.0761, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.28268060088157654, + "learning_rate": 2e-05, + "loss": 0.2757, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.523087739944458, + "learning_rate": 2e-05, + "loss": 0.1304, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.8083417415618896, + "learning_rate": 2e-05, + "loss": 0.5622, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.0939471498131752, + "learning_rate": 2e-05, + "loss": 0.0068, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.6428090929985046, + "learning_rate": 2e-05, + "loss": 0.1623, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.510278582572937, + "learning_rate": 2e-05, + "loss": 0.1102, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.2837804853916168, + "learning_rate": 2e-05, + "loss": 0.0153, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.910995602607727, + "learning_rate": 2e-05, + "loss": 0.0946, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.6980104446411133, + "learning_rate": 2e-05, + "loss": 0.2214, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.7017637491226196, + "learning_rate": 2e-05, + "loss": 0.0811, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.171759605407715, + "learning_rate": 2e-05, + "loss": 0.8276, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.7800589799880981, + "learning_rate": 2e-05, + "loss": 0.3885, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.41635552048683167, + "learning_rate": 2e-05, + "loss": 0.2662, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.4073396623134613, + "learning_rate": 2e-05, + "loss": 0.3289, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.489953875541687, + "learning_rate": 2e-05, + "loss": 0.082, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.2526872158050537, + "learning_rate": 2e-05, + "loss": 0.1934, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.5057745575904846, + "learning_rate": 2e-05, + "loss": 0.0336, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.8752731084823608, + "learning_rate": 2e-05, + "loss": 0.1089, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.235801935195923, + "learning_rate": 2e-05, + "loss": 0.1123, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.789495587348938, + "learning_rate": 2e-05, + "loss": 0.2236, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.4882419109344482, + "learning_rate": 2e-05, + "loss": 1.0651, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.8731786012649536, + "learning_rate": 2e-05, + "loss": 0.4021, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.099569082260132, + "learning_rate": 2e-05, + "loss": 0.3155, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.2141270786523819, + "learning_rate": 2e-05, + "loss": 0.0616, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291730211438592.0, + "train_loss": 0.2315644931793213, + "train_runtime": 131.511, + "train_samples_per_second": 3.042, + "train_steps_per_second": 0.76 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291730211438592.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..1619de04b07f6df028aecfdf01275ebfe8139616 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6e14e786e766d86cb4a79ca60156b5d6a9fe2e759be68adc1e9088d86a005ab9 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ed4da9967f5cb572d95239a683d4e697a3cfd3e0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6e5faa18694f6c2657a22983cf120dcf559de9d5a0f727ed0bfe407c182c0e7d +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..50f78bfa82434c50322a57259d0cc32497c088e7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fca5de2060711b070e59fadf5ec0a5807aaec4f3f6e272a54d4de5ce20a81460 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..1f295152ddb7db12e4b431893f0096d06f80c8a8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:51467579ab5666cff52b74911699787f46efee5069716d19393ad7085ecfd51c +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..0168b661956506905334bf57b06cc82c33bcbd88 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b148e11580ecb9edd8c35bd71d562019a568b78c4a05ee49de7230e4c7aafffc +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..60a831fa28bd65171492ec4498a9bd2adee87a8e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3127568f7bfbd54ef368cc06aa99f8380c6d7dc3e4ed94dfac7e6cb0e551a0b2 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..0378d82a8e1e7f61771d222b2a7b14c25d2b1e77 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cf858a8ce01e0d46c7d1deb719f85a864a74541fb5d3c6f10619c554ba4fd244 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..5871cf079113823b21a095fce6ca021afd7ef4bd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6e9cfd9c366a1edc71384707b177d555587b7546289ca22b600398f57dde01f6 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..4211f342fd2ed82523f36fb2dfd0b43b1c4d408f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.237710475921631, + "learning_rate": 2e-05, + "loss": 0.3775, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.075661659240723, + "learning_rate": 2e-05, + "loss": 0.1633, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.5205190181732178, + "learning_rate": 2e-05, + "loss": 0.1378, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.622267007827759, + "learning_rate": 2e-05, + "loss": 0.5881, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.9110074043273926, + "learning_rate": 2e-05, + "loss": 0.3236, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.1979341357946396, + "learning_rate": 2e-05, + "loss": 0.0171, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 5.861177921295166, + "learning_rate": 2e-05, + "loss": 0.412, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.6493712663650513, + "learning_rate": 2e-05, + "loss": 0.1411, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.6948139667510986, + "learning_rate": 2e-05, + "loss": 0.1241, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.6396384239196777, + "learning_rate": 2e-05, + "loss": 0.2959, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.4262688159942627, + "learning_rate": 2e-05, + "loss": 0.2082, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.7297395467758179, + "learning_rate": 2e-05, + "loss": 0.139, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.367156028747559, + "learning_rate": 2e-05, + "loss": 0.5541, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 8.937020301818848, + "learning_rate": 2e-05, + "loss": 0.6138, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 7.891814708709717, + "learning_rate": 2e-05, + "loss": 0.5288, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 4.096354961395264, + "learning_rate": 2e-05, + "loss": 0.0858, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 4.306092262268066, + "learning_rate": 2e-05, + "loss": 0.1612, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.2981480658054352, + "learning_rate": 2e-05, + "loss": 0.299, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.1092424392700195, + "learning_rate": 2e-05, + "loss": 0.1417, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.3104164600372314, + "learning_rate": 2e-05, + "loss": 0.3615, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.4578544497489929, + "learning_rate": 2e-05, + "loss": 0.0579, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.26196029782295227, + "learning_rate": 2e-05, + "loss": 0.6923, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.40582627058029175, + "learning_rate": 2e-05, + "loss": 0.0483, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.7194502949714661, + "learning_rate": 2e-05, + "loss": 0.2695, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.195401668548584, + "learning_rate": 2e-05, + "loss": 0.1189, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.150017023086548, + "learning_rate": 2e-05, + "loss": 0.9142, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 9.466475486755371, + "learning_rate": 2e-05, + "loss": 1.0176, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.3966981768608093, + "learning_rate": 2e-05, + "loss": 0.0192, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 5.472323417663574, + "learning_rate": 2e-05, + "loss": 0.1107, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.4477386474609375, + "learning_rate": 2e-05, + "loss": 0.4857, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 5.645395278930664, + "learning_rate": 2e-05, + "loss": 0.4622, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.9097378253936768, + "learning_rate": 2e-05, + "loss": 0.1038, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.425431966781616, + "learning_rate": 2e-05, + "loss": 0.7694, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.149527072906494, + "learning_rate": 2e-05, + "loss": 0.3325, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.333455204963684, + "learning_rate": 2e-05, + "loss": 0.0557, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 12.260011672973633, + "learning_rate": 2e-05, + "loss": 0.843, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.411311626434326, + "learning_rate": 2e-05, + "loss": 0.2913, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 8.59972858428955, + "learning_rate": 2e-05, + "loss": 0.3858, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 4.666265964508057, + "learning_rate": 2e-05, + "loss": 0.7488, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.537940740585327, + "learning_rate": 2e-05, + "loss": 0.6857, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 6.343166828155518, + "learning_rate": 2e-05, + "loss": 0.4238, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.07165875285863876, + "learning_rate": 2e-05, + "loss": 0.0074, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 5.735085964202881, + "learning_rate": 2e-05, + "loss": 0.3912, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 4.1733222007751465, + "learning_rate": 2e-05, + "loss": 0.4803, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.1532915830612183, + "learning_rate": 2e-05, + "loss": 0.1442, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.473877429962158, + "learning_rate": 2e-05, + "loss": 0.4735, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.018483638763428, + "learning_rate": 2e-05, + "loss": 0.6923, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.410689115524292, + "learning_rate": 2e-05, + "loss": 0.1287, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.6832528114318848, + "learning_rate": 2e-05, + "loss": 0.2981, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.4717618525028229, + "learning_rate": 2e-05, + "loss": 0.4562, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2221568423886848.0, + "train_loss": 0.351633186340332, + "train_runtime": 85.7428, + "train_samples_per_second": 4.665, + "train_steps_per_second": 1.166 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2221568423886848.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..bc25b0dc01fd3f24f90324f480bd8414f88a8644 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4696a836891fac7201d2b37a1b0724ac9cdfd7792f9d71e0433578157b233b6d +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..e866b453efa1af7424e4aba72bb932c8f38b2c52 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:96be612bab95da34c58a887d0852cafb8fe29713b1a20fdc4f5fd62af49e14cd +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..1f751021a828361497045f6cb5ff5fd5b7782aed --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56f1f1e2a580e6e66d2868df76e1138728b4b3ec9b3350c89535a4d42c51d23c +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..40794c399d8df64cd071616f3bf7727dccaecab7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fe020b8412b3e4e0416eafd71cc6158c0862958c66c5093d85ccb4fd6b0ddd56 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5e75226099c8168f6d6699d2515579dc69253135 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f4043fef56466786ce12a9bdc5f6af73c46b5925ad730c208ad1d63f504906e5 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..c1f8d946674db45c6fb60cf1ec123c9505ca4f9c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e3b6640e154c1485d1c17ae9f87d982575841ee7e47ab81ed85d568a30c96d09 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..ad901a2039564d43a1841a55d3975331bea4710e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:31173a2fcd5f24a37673b4124bb1c7ec371bb6f36f28e7c4379faf469fd6ddc8 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..af190370d0a0942356d7de7b83e7add19feba07f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c19376fc3d13e5593ee72ae219357dfa8c2343a23cfa328d3824a93f5e794adc +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..dd89af8f3d40fc88925111e83f1ec2b248609a66 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.7210659980773926, + "learning_rate": 2e-05, + "loss": 0.5298, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.909065008163452, + "learning_rate": 2e-05, + "loss": 0.4553, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.3670657873153687, + "learning_rate": 2e-05, + "loss": 0.3517, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.6238828897476196, + "learning_rate": 2e-05, + "loss": 0.3679, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.714184522628784, + "learning_rate": 2e-05, + "loss": 0.6914, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.143097996711731, + "learning_rate": 2e-05, + "loss": 0.4019, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.9192087650299072, + "learning_rate": 2e-05, + "loss": 0.917, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.800840139389038, + "learning_rate": 2e-05, + "loss": 0.7954, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.9944460988044739, + "learning_rate": 2e-05, + "loss": 0.5476, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.7866709232330322, + "learning_rate": 2e-05, + "loss": 0.3749, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.5968004465103149, + "learning_rate": 2e-05, + "loss": 0.2877, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.3698647022247314, + "learning_rate": 2e-05, + "loss": 0.4904, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.1019363403320312, + "learning_rate": 2e-05, + "loss": 0.4016, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.4402246475219727, + "learning_rate": 2e-05, + "loss": 0.3818, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.9486132860183716, + "learning_rate": 2e-05, + "loss": 0.4575, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.5378930568695068, + "learning_rate": 2e-05, + "loss": 0.3835, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.711106061935425, + "learning_rate": 2e-05, + "loss": 0.52, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.6947007179260254, + "learning_rate": 2e-05, + "loss": 0.4702, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.5925204157829285, + "learning_rate": 2e-05, + "loss": 0.5283, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.6147003173828125, + "learning_rate": 2e-05, + "loss": 0.438, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.779548168182373, + "learning_rate": 2e-05, + "loss": 0.3621, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.689056396484375, + "learning_rate": 2e-05, + "loss": 0.5988, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.0611324310302734, + "learning_rate": 2e-05, + "loss": 0.4958, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.4579930007457733, + "learning_rate": 2e-05, + "loss": 0.2568, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 3.1606550216674805, + "learning_rate": 2e-05, + "loss": 0.6178, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.103494644165039, + "learning_rate": 2e-05, + "loss": 0.5015, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.4368047714233398, + "learning_rate": 2e-05, + "loss": 0.3352, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.7359483242034912, + "learning_rate": 2e-05, + "loss": 0.2656, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.5112717151641846, + "learning_rate": 2e-05, + "loss": 0.2656, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.6594401597976685, + "learning_rate": 2e-05, + "loss": 0.313, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.4133126735687256, + "learning_rate": 2e-05, + "loss": 0.4596, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.1251816749572754, + "learning_rate": 2e-05, + "loss": 0.3715, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.1415674686431885, + "learning_rate": 2e-05, + "loss": 0.6284, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.19014789164066315, + "learning_rate": 2e-05, + "loss": 0.1744, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.586962342262268, + "learning_rate": 2e-05, + "loss": 0.33, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.947158932685852, + "learning_rate": 2e-05, + "loss": 0.4672, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.8931678533554077, + "learning_rate": 2e-05, + "loss": 0.6538, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.7501760721206665, + "learning_rate": 2e-05, + "loss": 0.3528, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 3.680293321609497, + "learning_rate": 2e-05, + "loss": 0.6068, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.9687423706054688, + "learning_rate": 2e-05, + "loss": 0.6504, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.5083470344543457, + "learning_rate": 2e-05, + "loss": 0.752, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.6850969791412354, + "learning_rate": 2e-05, + "loss": 0.708, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.5290727615356445, + "learning_rate": 2e-05, + "loss": 0.3638, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.534722924232483, + "learning_rate": 2e-05, + "loss": 0.4766, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.012035796418786049, + "learning_rate": 2e-05, + "loss": 0.502, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.7276091575622559, + "learning_rate": 2e-05, + "loss": 0.345, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.1122796535491943, + "learning_rate": 2e-05, + "loss": 0.4854, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.7822918891906738, + "learning_rate": 2e-05, + "loss": 0.4297, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.3838807344436646, + "learning_rate": 2e-05, + "loss": 0.3823, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.260063886642456, + "learning_rate": 2e-05, + "loss": 0.4541, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2192326831112192.0, + "train_loss": 0.4679565715789795, + "train_runtime": 77.3995, + "train_samples_per_second": 5.168, + "train_steps_per_second": 1.292 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2192326831112192.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..1c5975a26a5618ba3f244877bc7826178d8502ab --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:95f0d8ed12d6463e3a164d2127e9c8a37107f0a24343c00e3f45b95792615dfb +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..1ff5bf885331252b3404d7e97c62b18a625f360b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8ff6593c1e87ee157c6e13a17723b186fa803a174c51d9ffa547bc3a47dcbc4b +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..9480522b018e7d4903014535d94c4ed258474bda --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ac8880e1bbe5caddd15b471aae216a25d40a42e159ee99d60e07f53bd84acd49 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..6cb7f753a31c8535d32e4aa70083655a8a800801 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4e82b381c4b2cb1ce469b8e29cee5ceeafcc0e682fb94a0bac8093f823e49a51 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..a76e7d7a9f0ee3548cdc230dc2a0f56326f7f4af --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5237a861d053df735cc8376826f1b392fdfb90964395683a121475ac673e61f4 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..0cb65224075e372f2735ac0ff75c62d2450057e3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0d2d5047a59c6fbf01c01dfc00b3c968f56372bc47b3497e11ac1ac481d61c33 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..cd5f145df31e3e8f3f4653d1bd3e4a3ff3385eac --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e3664c7fbddd707ff08bd2db8e38e12571fc315f9310de2021d0c3c097b7dba3 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..b32d4a71518e05ad774336fc768a5abbeeb7ddcc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7dbb1ee0a729e122da0e64a0d4f9a4fa8321c60135775d673b20c73c8e84bb7c +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..81dc374d59ab9f8463aa38a55a8cf7c22b7346b3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.9738138914108276, + "learning_rate": 2e-05, + "loss": 0.0441, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.0016529003623872995, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.061222709715366364, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.5949037075042725, + "learning_rate": 2e-05, + "loss": 0.0277, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.21504652500152588, + "learning_rate": 2e-05, + "loss": 0.0147, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.06401089578866959, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.1229968070983887, + "learning_rate": 2e-05, + "loss": 0.6408, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.197314977645874, + "learning_rate": 2e-05, + "loss": 0.2815, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.24723775684833527, + "learning_rate": 2e-05, + "loss": 0.1509, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.057056427001953, + "learning_rate": 2e-05, + "loss": 0.1423, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.09781523793935776, + "learning_rate": 2e-05, + "loss": 0.1219, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.1981300413608551, + "learning_rate": 2e-05, + "loss": 0.0152, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.28713059425354004, + "learning_rate": 2e-05, + "loss": 0.0434, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.6210917234420776, + "learning_rate": 2e-05, + "loss": 0.1274, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.020366854965686798, + "learning_rate": 2e-05, + "loss": 0.0117, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.055765021592378616, + "learning_rate": 2e-05, + "loss": 0.0116, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.08080822974443436, + "learning_rate": 2e-05, + "loss": 0.0338, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.11208204180002213, + "learning_rate": 2e-05, + "loss": 0.0739, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.12620840966701508, + "learning_rate": 2e-05, + "loss": 0.0239, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.4830634891986847, + "learning_rate": 2e-05, + "loss": 0.0437, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.4966158866882324, + "learning_rate": 2e-05, + "loss": 0.3203, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.025074919685721397, + "learning_rate": 2e-05, + "loss": 0.0328, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.215610384941101, + "learning_rate": 2e-05, + "loss": 0.1837, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.5805869102478027, + "learning_rate": 2e-05, + "loss": 0.0299, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.040374137461185455, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.06280364841222763, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.032653067260980606, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.08511438220739365, + "learning_rate": 2e-05, + "loss": 0.0036, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.017775876447558403, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.057130321860313416, + "learning_rate": 2e-05, + "loss": 0.0039, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.008392478339374065, + "learning_rate": 2e-05, + "loss": 0.0077, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.010977689176797867, + "learning_rate": 2e-05, + "loss": 0.0047, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.10285238176584244, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.6860859394073486, + "learning_rate": 2e-05, + "loss": 0.0311, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.006630100309848785, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.022569384425878525, + "learning_rate": 2e-05, + "loss": 0.2734, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.03492208942770958, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.006131832022219896, + "learning_rate": 2e-05, + "loss": 0.0182, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.3356422185897827, + "learning_rate": 2e-05, + "loss": 0.1706, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.0036606271751224995, + "learning_rate": 2e-05, + "loss": 0.0099, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.03982897475361824, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.04343670234084129, + "learning_rate": 2e-05, + "loss": 0.0037, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0027983614709228277, + "learning_rate": 2e-05, + "loss": 0.0086, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.5703356266021729, + "learning_rate": 2e-05, + "loss": 0.1706, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.003599955700337887, + "learning_rate": 2e-05, + "loss": 0.0146, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.27015066146850586, + "learning_rate": 2e-05, + "loss": 0.0714, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.334803342819214, + "learning_rate": 2e-05, + "loss": 0.4228, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.026150742545723915, + "learning_rate": 2e-05, + "loss": 0.4527, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.03261609375476837, + "learning_rate": 2e-05, + "loss": 0.0021, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.0613900423049927, + "learning_rate": 2e-05, + "loss": 0.035, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5287535144075264.0, + "train_loss": 0.08226008832454682, + "train_runtime": 127.7692, + "train_samples_per_second": 3.131, + "train_steps_per_second": 0.783 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5287535144075264.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..1844b26a2299bd624a6c048eb1203e8fde82ba5d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:35ece1dca8be3b6685b71db818d334d405b3274b312d231d209d44aa565498a3 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..7b0a0a06dd419d444ee044d3b5ce58b4dd0be1d7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7363d749b5897d0a0106afd33682d1b5360f47bc823196a6febed834d62bc5e2 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..2edc5a56804cb93a8fb50499f1d24db6368b1a7d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:65e60b2f9b43f48ebe79115d1af022f494884cf62b2561b63eff2911c7eb2745 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0bca64fd7c5c2e3a040a071a1c3d05cc50cf85d2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ad177f0c3bd80c3efbfe73a3350e31b9300e8dfd5a728aa3fb5c0f09aca9d4f5 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..392bc234bc855fcef66b3fa2273773182054f21a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:11d5e1acddb26ab7a324e25284e5249e0758253eb40a44e9c2122f1b00751a8d +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..71ea66e1b749e4d392fde5013e880dbcbdc9e4e9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d58fb7378815b9f346f4bf2caf386b4af2c00a00e93a57b0395258df04f4125c +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..00580dce192a45cc91b8276b8beed2d4a144a9de --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e0525ee20f2d75eec6813bd3c05b925f423c68b295f65106355b84acd4a0dcd4 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..7b7180cd3990c8384f85e9bfcc0ed6d0f3c6affd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:31057076c8da1b26c28440081f8830740fc7efb050487fbbc88c9151f53cf471 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3af00e52d2aa61a518eb33bd0f919baa3b91c06e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.6773409843444824, + "learning_rate": 2e-05, + "loss": 0.2477, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.3861768245697021, + "learning_rate": 2e-05, + "loss": 0.2258, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.8264509439468384, + "learning_rate": 2e-05, + "loss": 0.4564, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.8123915195465088, + "learning_rate": 2e-05, + "loss": 0.0823, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.6776534914970398, + "learning_rate": 2e-05, + "loss": 0.1022, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.16754977405071259, + "learning_rate": 2e-05, + "loss": 0.0361, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.5837736129760742, + "learning_rate": 2e-05, + "loss": 0.0855, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.648746967315674, + "learning_rate": 2e-05, + "loss": 0.3412, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.817319631576538, + "learning_rate": 2e-05, + "loss": 0.2317, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.9044545292854309, + "learning_rate": 2e-05, + "loss": 1.2299, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.748376846313477, + "learning_rate": 2e-05, + "loss": 0.6151, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 5.7237348556518555, + "learning_rate": 2e-05, + "loss": 0.2781, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.9041240215301514, + "learning_rate": 2e-05, + "loss": 1.0011, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.3380802869796753, + "learning_rate": 2e-05, + "loss": 0.1512, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.3328051567077637, + "learning_rate": 2e-05, + "loss": 0.2772, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.1047422885894775, + "learning_rate": 2e-05, + "loss": 0.1303, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5501200556755066, + "learning_rate": 2e-05, + "loss": 0.2404, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.6180824041366577, + "learning_rate": 2e-05, + "loss": 0.2881, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.1546880006790161, + "learning_rate": 2e-05, + "loss": 0.1126, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.454337239265442, + "learning_rate": 2e-05, + "loss": 0.2541, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.045667678117752075, + "learning_rate": 2e-05, + "loss": 0.2411, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.9230005741119385, + "learning_rate": 2e-05, + "loss": 0.3701, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.02381601184606552, + "learning_rate": 2e-05, + "loss": 0.0136, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.4391237497329712, + "learning_rate": 2e-05, + "loss": 0.2138, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.17225654423236847, + "learning_rate": 2e-05, + "loss": 0.1779, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.603492259979248, + "learning_rate": 2e-05, + "loss": 0.3031, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.4573830366134644, + "learning_rate": 2e-05, + "loss": 0.3351, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.9391083717346191, + "learning_rate": 2e-05, + "loss": 0.1601, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.059376079589128494, + "learning_rate": 2e-05, + "loss": 0.0465, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.12897945940494537, + "learning_rate": 2e-05, + "loss": 0.0819, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.0584347248077393, + "learning_rate": 2e-05, + "loss": 0.2339, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.3905250132083893, + "learning_rate": 2e-05, + "loss": 0.0899, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.2331008911132812, + "learning_rate": 2e-05, + "loss": 0.4036, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.2745652198791504, + "learning_rate": 2e-05, + "loss": 0.2753, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.1008955240249634, + "learning_rate": 2e-05, + "loss": 0.22, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.26280224323272705, + "learning_rate": 2e-05, + "loss": 0.0461, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.4440811276435852, + "learning_rate": 2e-05, + "loss": 0.5445, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.63106769323349, + "learning_rate": 2e-05, + "loss": 0.1125, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.0053333044052124, + "learning_rate": 2e-05, + "loss": 0.229, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.13429856300354, + "learning_rate": 2e-05, + "loss": 0.0954, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.7349315881729126, + "learning_rate": 2e-05, + "loss": 0.1144, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.413107395172119, + "learning_rate": 2e-05, + "loss": 0.7592, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.094510078430176, + "learning_rate": 2e-05, + "loss": 0.5425, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.571529746055603, + "learning_rate": 2e-05, + "loss": 0.1686, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.0386180877685547, + "learning_rate": 2e-05, + "loss": 0.1113, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.3208231925964355, + "learning_rate": 2e-05, + "loss": 0.5395, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.27220144867897034, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.6989845037460327, + "learning_rate": 2e-05, + "loss": 0.247, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.3015912175178528, + "learning_rate": 2e-05, + "loss": 0.316, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.37014102935791, + "learning_rate": 2e-05, + "loss": 0.2827, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5324358893436928.0, + "train_loss": 0.27346144676208495, + "train_runtime": 122.3964, + "train_samples_per_second": 3.268, + "train_steps_per_second": 0.817 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5324358893436928.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..979a34525c3a746c7fa0d458e823f92f078aa59f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c80abf1fa02b23c13d5e59794a71c143f06e42544b80323851f5bef9fa789179 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ee46de2213d5c48bffc0efd69eaa8ebac7f338b7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:febc9a14a768bd6010a4946522599e8393e32a116bc5610ec25e50d278ca4c25 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..0a0a8b5bce94c24af9d7c5c8f843d83adad308f6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb9e2eedf8b6dac3dcb02e3e787a8d270cda242fb32616dfaedc8abd1b561e96 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..f40c9ef73992c81b6552ac68562bb62a6f284493 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c95f2b5f1de313c1c06ca1583c32f0a7d213f68f423d09a638fa2cfa1fab5f74 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..6b058314f25607b8b7d261003c0ad80c1a6a8e18 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5670cdb6345754a7fc6eb909320849a1906677f0b5535472c8cde8448c642b8e +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..396c17f4deeb0ecd618c263dc2fb0adbcad4d7c4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:42130967b04d2b2ab3173c15390a807c6abee1d56a394c591a4921dd6dfdffca +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..40c6dbb2a5e9808c8800a22aac143ea1d0207004 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8b167ae60521653ae55dd7dfca7a6bffaa4b53bfc10a56868d5ba852e04ca02 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..3fec2ca1ecc17f690199f7ace2f78716ec228aad --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2291f76663d7754c4c91c5f75540c9070a92fc235d0a54b08960ddbca0fe495a +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..04048ee5dab4ecd1b82ceeb6b0d652b5b401913c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.6315317153930664, + "learning_rate": 2e-05, + "loss": 0.4695, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.9247969388961792, + "learning_rate": 2e-05, + "loss": 0.3093, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.20382171869277954, + "learning_rate": 2e-05, + "loss": 0.0538, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.06592469662427902, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.536984920501709, + "learning_rate": 2e-05, + "loss": 0.6108, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.4564790725708008, + "learning_rate": 2e-05, + "loss": 0.2434, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.522944211959839, + "learning_rate": 2e-05, + "loss": 0.3451, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.023859336972236633, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.0975253582000732, + "learning_rate": 2e-05, + "loss": 0.1319, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.4118559658527374, + "learning_rate": 2e-05, + "loss": 0.0294, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.19393041729927063, + "learning_rate": 2e-05, + "loss": 0.0307, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.02002054452896118, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.868740081787109, + "learning_rate": 2e-05, + "loss": 0.6108, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 6.986715793609619, + "learning_rate": 2e-05, + "loss": 0.3558, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.190476655960083, + "learning_rate": 2e-05, + "loss": 0.1999, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.7140227556228638, + "learning_rate": 2e-05, + "loss": 0.0634, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.19780878722667694, + "learning_rate": 2e-05, + "loss": 0.1912, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.6610605716705322, + "learning_rate": 2e-05, + "loss": 0.0384, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.15899765491485596, + "learning_rate": 2e-05, + "loss": 0.0158, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.46223783493042, + "learning_rate": 2e-05, + "loss": 0.3548, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.2458622306585312, + "learning_rate": 2e-05, + "loss": 0.1585, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.6165256500244141, + "learning_rate": 2e-05, + "loss": 0.3948, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.40964561700820923, + "learning_rate": 2e-05, + "loss": 0.0346, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.8917354345321655, + "learning_rate": 2e-05, + "loss": 0.05, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.02671065367758274, + "learning_rate": 2e-05, + "loss": 0.1612, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.026933759450912476, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.3557865619659424, + "learning_rate": 2e-05, + "loss": 0.0459, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.6702917814254761, + "learning_rate": 2e-05, + "loss": 0.1107, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.8045272827148438, + "learning_rate": 2e-05, + "loss": 0.185, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.016062943264842033, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7252837419509888, + "learning_rate": 2e-05, + "loss": 0.2136, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.3930991291999817, + "learning_rate": 2e-05, + "loss": 0.0266, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.008084774017334, + "learning_rate": 2e-05, + "loss": 0.1249, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.7267833948135376, + "learning_rate": 2e-05, + "loss": 0.2188, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.3362174332141876, + "learning_rate": 2e-05, + "loss": 0.1633, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.16612136363983154, + "learning_rate": 2e-05, + "loss": 0.0347, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.539078712463379, + "learning_rate": 2e-05, + "loss": 0.1265, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.407500743865967, + "learning_rate": 2e-05, + "loss": 0.154, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.7071353793144226, + "learning_rate": 2e-05, + "loss": 0.1072, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.2380239963531494, + "learning_rate": 2e-05, + "loss": 0.0804, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.42217519879341125, + "learning_rate": 2e-05, + "loss": 0.0552, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.7090271711349487, + "learning_rate": 2e-05, + "loss": 0.0573, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.7171026468276978, + "learning_rate": 2e-05, + "loss": 0.0262, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.2582712769508362, + "learning_rate": 2e-05, + "loss": 0.0226, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.6299818754196167, + "learning_rate": 2e-05, + "loss": 0.0701, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.428096055984497, + "learning_rate": 2e-05, + "loss": 0.0358, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.11565401405096054, + "learning_rate": 2e-05, + "loss": 0.0132, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.611966609954834, + "learning_rate": 2e-05, + "loss": 0.0369, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.24849525094032288, + "learning_rate": 2e-05, + "loss": 0.0144, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.003970960155129433, + "learning_rate": 2e-05, + "loss": 0.027, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5339656207990784.0, + "train_loss": 0.13631507158279418, + "train_runtime": 122.7468, + "train_samples_per_second": 3.259, + "train_steps_per_second": 0.815 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5339656207990784.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..e3d9d5cfedadf4fdaab720b2d4b864f42f8f1b72 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:407d22b4e4c660b6f68e8c79aee821f7007b58f9353ad00c80174ea4c3c343df +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..d079a62795043863b58210c51433f1089fb80dde --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2f21fd6686a3bc9e98f326ddc95985daeae271cffdb300c4c44778c5f1c6b3f2 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..b2e2ae58ec68e54026f912ad620b5f0d179983bf --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b114f488612250ac54aecb38f76fa0ec0afb236571fd651c41c58c2b7963f11 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..900e1642dbcc24985b0d1fefe847a4f1772de454 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:75c2c5c6a9dcc07dcfc942d0724baa2288314ef899bb1372d1b9d82bb4878f92 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..7bfbb882a1fb93562b20962518cfb013aaafd62f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7fd28c3dbfe997821358d4e820e5a4db3d90155fe4d123de393543e04163dd9f +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b70f37aafd13e056a7840d64d7bca9862eecb89c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b48abb22e70035a338d2bc5223ec2be06529c5c75ac9c128fd310378fafed124 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b7a45a9a8ab6c6da27d3750d796a2347eb34c992 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:971b680ccaa5fa60a3217f9455b074be0b7f88bfcf16ed132f933dcc4ca90923 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..261e62fceef9ebb78ad57714ca1565aafdec4fc6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1bf0a8ea3268029662921b0fa787b57ae08195b5979f7b875abc13e4bcce7354 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3d05a616ca4e2902210620b6cbac9cd266139048 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.377012252807617, + "learning_rate": 2e-05, + "loss": 0.1418, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.392651081085205, + "learning_rate": 2e-05, + "loss": 0.385, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 4.012504577636719, + "learning_rate": 2e-05, + "loss": 0.3775, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.9318208694458008, + "learning_rate": 2e-05, + "loss": 0.0784, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.6181113719940186, + "learning_rate": 2e-05, + "loss": 0.0923, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.5494726300239563, + "learning_rate": 2e-05, + "loss": 0.1821, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.6615825891494751, + "learning_rate": 2e-05, + "loss": 0.039, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.44751736521720886, + "learning_rate": 2e-05, + "loss": 0.3047, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.338043451309204, + "learning_rate": 2e-05, + "loss": 0.0765, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.148883581161499, + "learning_rate": 2e-05, + "loss": 0.1953, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.727282524108887, + "learning_rate": 2e-05, + "loss": 0.4316, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.823374032974243, + "learning_rate": 2e-05, + "loss": 0.2227, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.0644009113311768, + "learning_rate": 2e-05, + "loss": 0.3637, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.5839422941207886, + "learning_rate": 2e-05, + "loss": 0.2475, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.483096718788147, + "learning_rate": 2e-05, + "loss": 0.2003, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.617430567741394, + "learning_rate": 2e-05, + "loss": 0.0645, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.766664743423462, + "learning_rate": 2e-05, + "loss": 0.4334, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.13144594430923462, + "learning_rate": 2e-05, + "loss": 0.0436, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.218085527420044, + "learning_rate": 2e-05, + "loss": 0.291, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.8389110565185547, + "learning_rate": 2e-05, + "loss": 0.1438, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.18799269199371338, + "learning_rate": 2e-05, + "loss": 0.058, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.05910325050354, + "learning_rate": 2e-05, + "loss": 0.5908, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.9163956046104431, + "learning_rate": 2e-05, + "loss": 0.0408, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.7585868835449219, + "learning_rate": 2e-05, + "loss": 0.0701, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.032984733581543, + "learning_rate": 2e-05, + "loss": 0.6447, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.017808437347412, + "learning_rate": 2e-05, + "loss": 0.2701, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.1482895463705063, + "learning_rate": 2e-05, + "loss": 0.0135, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.3856992721557617, + "learning_rate": 2e-05, + "loss": 0.0275, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.729601502418518, + "learning_rate": 2e-05, + "loss": 0.2009, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5952390432357788, + "learning_rate": 2e-05, + "loss": 0.0952, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 8.146366119384766, + "learning_rate": 2e-05, + "loss": 0.5844, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.3974289894104004, + "learning_rate": 2e-05, + "loss": 0.1243, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.0182135105133057, + "learning_rate": 2e-05, + "loss": 0.1752, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.594396591186523, + "learning_rate": 2e-05, + "loss": 0.4948, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.4696035087108612, + "learning_rate": 2e-05, + "loss": 0.3777, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.294278860092163, + "learning_rate": 2e-05, + "loss": 0.2165, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.46589335799217224, + "learning_rate": 2e-05, + "loss": 0.1661, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.4711010456085205, + "learning_rate": 2e-05, + "loss": 0.3441, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.9517972469329834, + "learning_rate": 2e-05, + "loss": 0.0456, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.189296841621399, + "learning_rate": 2e-05, + "loss": 0.1635, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.541762113571167, + "learning_rate": 2e-05, + "loss": 0.1599, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.291274070739746, + "learning_rate": 2e-05, + "loss": 0.1752, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.820042371749878, + "learning_rate": 2e-05, + "loss": 0.5785, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.68217658996582, + "learning_rate": 2e-05, + "loss": 0.6943, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.011737823486328, + "learning_rate": 2e-05, + "loss": 0.2762, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.33684635162353516, + "learning_rate": 2e-05, + "loss": 0.1801, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.4390194416046143, + "learning_rate": 2e-05, + "loss": 0.0443, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.37847065925598145, + "learning_rate": 2e-05, + "loss": 0.7382, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.3042492866516113, + "learning_rate": 2e-05, + "loss": 0.2653, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.3256466388702393, + "learning_rate": 2e-05, + "loss": 0.0884, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2214526502043648.0, + "train_loss": 0.244374361038208, + "train_runtime": 80.6155, + "train_samples_per_second": 4.962, + "train_steps_per_second": 1.24 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2214526502043648.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..cbf36bbf96f2756100f60f16b01c0127ee4fd64e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1adfc8950c819b06540115fbc6a75bf7893e6e682e377992501083dcfa1be15a +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..698668778fa9db755eecbc54a6851eacd57e586b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7aee5d721983869b527c2f493379d33b0beff12a64de35f85d0202c5107405e7 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..526957f75da58a6aaaca07dc5ff2fa190f74da91 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a39e96b721ec832db87be7e4ba2caba3a7d5d4c11701de3c4107f7242099a4a9 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d481f5e407ec6b4a45ae7e0c7c14ae4d788cf218 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ffd7bacdb7448c904c6a4b43bbf24ab813b86d021ebe39bff8de40e54ed9661b +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..118613e5b49b7c1f10b49eaaf9451ea4f8dbebd0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fd6750eab8f2c4f6a26aafbedce1005b5072fb9b1fcd274356ec34bf5121d7bd +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..f6bac445a83f8ceb650fb226151e45901ebadd84 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:00ff4d08d9a8dbfe2a3e27442fd1978c6654616d24fa2e521f564dd732b8fce1 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..556ea90f21c29683f30a6bcff350bab207bbb004 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f19f059b52e7d84054efd1f7982670cfcf96a5a14bd6a18cdc1630ccb5ce1cbf +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d65c9567e1bb76ea2ca237838fbb66359693809d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:957a8a93a27929fcd0ba01a7ca80612123534a86f12b9c6224d0c38f437a9ade +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..aa2d5486eb6b4e76043b0496c648d74ed3b7a488 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.1274750232696533, + "learning_rate": 2e-05, + "loss": 0.2684, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.24790239334106445, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 7.876332759857178, + "learning_rate": 2e-05, + "loss": 0.8587, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.4195573329925537, + "learning_rate": 2e-05, + "loss": 0.1692, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.0443975925445557, + "learning_rate": 2e-05, + "loss": 0.1748, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.10934049636125565, + "learning_rate": 2e-05, + "loss": 0.0233, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.8561116456985474, + "learning_rate": 2e-05, + "loss": 0.2069, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.4809703826904297, + "learning_rate": 2e-05, + "loss": 0.3133, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.1467329263687134, + "learning_rate": 2e-05, + "loss": 0.0564, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.06186947971582413, + "learning_rate": 2e-05, + "loss": 0.0219, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.7323760986328125, + "learning_rate": 2e-05, + "loss": 0.5676, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3583230972290039, + "learning_rate": 2e-05, + "loss": 0.1591, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.3560287952423096, + "learning_rate": 2e-05, + "loss": 0.4785, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.9718797206878662, + "learning_rate": 2e-05, + "loss": 0.2046, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.0406503677368164, + "learning_rate": 2e-05, + "loss": 0.5583, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.1780941486358643, + "learning_rate": 2e-05, + "loss": 0.2299, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.38859978318214417, + "learning_rate": 2e-05, + "loss": 0.2426, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.243556499481201, + "learning_rate": 2e-05, + "loss": 0.2425, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.9962045550346375, + "learning_rate": 2e-05, + "loss": 0.0748, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.5472426414489746, + "learning_rate": 2e-05, + "loss": 0.2607, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.09325888007879257, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.36677417159080505, + "learning_rate": 2e-05, + "loss": 0.0989, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.569000005722046, + "learning_rate": 2e-05, + "loss": 0.089, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.3031120300292969, + "learning_rate": 2e-05, + "loss": 0.088, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.256440132856369, + "learning_rate": 2e-05, + "loss": 0.0346, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.3505374193191528, + "learning_rate": 2e-05, + "loss": 0.0433, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.1625138521194458, + "learning_rate": 2e-05, + "loss": 0.0553, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 4.943670272827148, + "learning_rate": 2e-05, + "loss": 0.3898, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.7390772104263306, + "learning_rate": 2e-05, + "loss": 0.2349, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.6583502888679504, + "learning_rate": 2e-05, + "loss": 0.0588, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.33168643712997437, + "learning_rate": 2e-05, + "loss": 0.0159, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.818939208984375, + "learning_rate": 2e-05, + "loss": 0.4669, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.2423162460327148, + "learning_rate": 2e-05, + "loss": 0.0874, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 8.29089069366455, + "learning_rate": 2e-05, + "loss": 1.4476, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.6976922750473022, + "learning_rate": 2e-05, + "loss": 0.2002, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 5.96759557723999, + "learning_rate": 2e-05, + "loss": 0.3339, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.6955366134643555, + "learning_rate": 2e-05, + "loss": 0.7075, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.03233902528882027, + "learning_rate": 2e-05, + "loss": 0.0033, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.0800695419311523, + "learning_rate": 2e-05, + "loss": 0.0522, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.3585282862186432, + "learning_rate": 2e-05, + "loss": 0.0418, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.4090957641601562, + "learning_rate": 2e-05, + "loss": 0.4405, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.8466939330101013, + "learning_rate": 2e-05, + "loss": 0.2204, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.9248936772346497, + "learning_rate": 2e-05, + "loss": 0.0265, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.341991901397705, + "learning_rate": 2e-05, + "loss": 0.2365, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.6409504413604736, + "learning_rate": 2e-05, + "loss": 0.3578, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.245504856109619, + "learning_rate": 2e-05, + "loss": 0.1159, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.507658958435059, + "learning_rate": 2e-05, + "loss": 0.7485, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 4.716681957244873, + "learning_rate": 2e-05, + "loss": 0.1511, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.4564982652664185, + "learning_rate": 2e-05, + "loss": 0.094, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.1849710941314697, + "learning_rate": 2e-05, + "loss": 0.1405, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2207748649385984.0, + "train_loss": 0.24261905670166015, + "train_runtime": 78.2776, + "train_samples_per_second": 5.11, + "train_steps_per_second": 1.278 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2207748649385984.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..00930ba74900819577a5f0c0cefbd790883098e2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b414b583204fce9d35b8d6eb26f5c77f9b49f7fedf333dfd8b82e8b934cc0418 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..89a6e708af9b975dcdf5359e6c8bab69073e4b59 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bdc02013c31ec71a3ad7bc0d0898b87ecd055d18b1f4aa9e2eeb7359e78545ac +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..098755369254e8acd7369c70b227279fdf491040 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:58ad2218854aa0f08d39d761b02365db2f681486cdfb8850411bb26169703b0a +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..32e6e839fc84515dfcbe6365f839cc2f8c2a5e9d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:513f920b83f6c134d6a36adbf72bda0e68c40518112d3b77ca292af524b66507 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..8e209802e1d6091296b6443748f2f791590405a3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8de08446c85df92bbf597fea2a05486d6c04969b47f23ac17c961cd2e005320f +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b98b0786db5f831ac9dc0a47789b8fa26a9fde38 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:61d315c2cd6a0e023dd6af7cfc3a29b5d1364a7df7dea6d333a9c702f9112195 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..ede02fa7bcb86958c14fd7c223b8f51a3d3b3fd7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:06715dafff5593311f7568b7f09b07677d0111b0178ce8e96304b0b3b5de913c +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..deebcdfe0ad1945c5a59ccd7f4ed4c893e72a6aa --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:330c87648fe0ee98a1ad9b9cab2e057a22743ba63b9450845ba5dde0afbaf2b6 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..355e677e95d19093f4b65aabc49666b67bc32b1a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.848058819770813, + "learning_rate": 2e-05, + "loss": 0.0626, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.7377076148986816, + "learning_rate": 2e-05, + "loss": 0.4752, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.0862412452697754, + "learning_rate": 2e-05, + "loss": 0.1135, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.8021219372749329, + "learning_rate": 2e-05, + "loss": 0.0553, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.8433351516723633, + "learning_rate": 2e-05, + "loss": 0.1238, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.1266226768493652, + "learning_rate": 2e-05, + "loss": 0.1706, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.6890785098075867, + "learning_rate": 2e-05, + "loss": 0.0241, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.099366188049316, + "learning_rate": 2e-05, + "loss": 0.2369, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 6.179292678833008, + "learning_rate": 2e-05, + "loss": 0.3326, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 6.927247524261475, + "learning_rate": 2e-05, + "loss": 0.9438, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.2217345237731934, + "learning_rate": 2e-05, + "loss": 0.1773, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 6.076210975646973, + "learning_rate": 2e-05, + "loss": 0.7074, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.0832374095916748, + "learning_rate": 2e-05, + "loss": 0.0887, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.13343863189220428, + "learning_rate": 2e-05, + "loss": 0.0078, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.6522243022918701, + "learning_rate": 2e-05, + "loss": 0.1447, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.2924253940582275, + "learning_rate": 2e-05, + "loss": 0.109, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.823907732963562, + "learning_rate": 2e-05, + "loss": 0.0626, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.35624608397483826, + "learning_rate": 2e-05, + "loss": 0.0173, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.6828482151031494, + "learning_rate": 2e-05, + "loss": 0.2667, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.4852460622787476, + "learning_rate": 2e-05, + "loss": 0.2579, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.35870397090911865, + "learning_rate": 2e-05, + "loss": 0.0217, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.0098726749420166, + "learning_rate": 2e-05, + "loss": 0.3021, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.305283784866333, + "learning_rate": 2e-05, + "loss": 0.0458, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 4.0358147621154785, + "learning_rate": 2e-05, + "loss": 0.225, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.8295071125030518, + "learning_rate": 2e-05, + "loss": 0.3641, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.970831036567688, + "learning_rate": 2e-05, + "loss": 0.1545, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.6733829975128174, + "learning_rate": 2e-05, + "loss": 0.1818, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 5.25558614730835, + "learning_rate": 2e-05, + "loss": 0.2789, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 5.518392086029053, + "learning_rate": 2e-05, + "loss": 0.3291, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 6.368981838226318, + "learning_rate": 2e-05, + "loss": 0.4112, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 6.831753253936768, + "learning_rate": 2e-05, + "loss": 0.9485, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.354714870452881, + "learning_rate": 2e-05, + "loss": 0.8714, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.3441479802131653, + "learning_rate": 2e-05, + "loss": 0.0598, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.936461448669434, + "learning_rate": 2e-05, + "loss": 0.3112, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.615408182144165, + "learning_rate": 2e-05, + "loss": 0.0401, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.368828058242798, + "learning_rate": 2e-05, + "loss": 0.489, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.7747846841812134, + "learning_rate": 2e-05, + "loss": 0.2931, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.4268436431884766, + "learning_rate": 2e-05, + "loss": 0.661, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.181837320327759, + "learning_rate": 2e-05, + "loss": 0.1732, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.3047047257423401, + "learning_rate": 2e-05, + "loss": 0.1259, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.8474886417388916, + "learning_rate": 2e-05, + "loss": 0.2108, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.792332172393799, + "learning_rate": 2e-05, + "loss": 0.4102, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.9686474800109863, + "learning_rate": 2e-05, + "loss": 0.4203, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.1518681049346924, + "learning_rate": 2e-05, + "loss": 0.042, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.749695301055908, + "learning_rate": 2e-05, + "loss": 0.541, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.9528850317001343, + "learning_rate": 2e-05, + "loss": 0.152, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.08033628016710281, + "learning_rate": 2e-05, + "loss": 0.0459, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.8656932711601257, + "learning_rate": 2e-05, + "loss": 0.0471, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.031501177698373795, + "learning_rate": 2e-05, + "loss": 0.0808, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.0148407220840454, + "learning_rate": 2e-05, + "loss": 0.0557, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2212818824724480.0, + "train_loss": 0.2534214687347412, + "train_runtime": 76.6849, + "train_samples_per_second": 5.216, + "train_steps_per_second": 1.304 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2212818824724480.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..42839a6f29deaaa730a6c846ac2488a518b30c49 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:84445c2e052bc5da65f2500b1bb52bfde63c5cac9ad85fe126dadf355855bd56 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..7e4bd0b4e3e6339f8f7bc3523f6e7a7a74fd2433 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:952bda9c8102e39384f4001a2ec8ed80274aa82fc2df65eea29d7aaddaa525ac +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..11d5f4e17c2cb56b0ec45c1df89d219418090236 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:468bd96cf63d0ddad318ff4d30bf08501c74a227a0bc87cc955eaeee316e0be8 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..36b0240e97ae51ac0346577996439e49ba994e31 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0495a83b57b15ea3d45f2231a2301ae0d8c4e305b3499ca0145af9b1a025f3d9 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..0f8aaf3351c2c7c45f82022f60be2c0bee2c3286 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:60c60b1e47066097bd7a71d8c257c71d1b660ecae1668e5a1183df1bddc9df4b +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..2f2510e68cec82e9577a5dfea866ce7038a8d2e0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:95095956fbabba3d3380371b7ecdaf4187c7af959c81e0fd3d4578d16416a345 +size 369839594 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..73064fc4cc708daa22349e774a9371df23427d66 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a4d1e112001a88d61fbb92fa6389f88683cc0d2120fe8865b9eeaf9d74c9e4eb +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..c696c0d6caaa09822169065e8bf1a002109742ab --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9076e709cb582258c6625f57a2dfd72aca698793519a0099cddd5721fc092017 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..f31dd60629f5454b990a4a6285a316c33b5b56fb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.08075618743896484, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.0558048486709595, + "learning_rate": 2e-05, + "loss": 0.0618, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.48747897148132324, + "learning_rate": 2e-05, + "loss": 0.0153, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.03438630327582359, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.06621252745389938, + "learning_rate": 2e-05, + "loss": 0.3849, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 5.463819980621338, + "learning_rate": 2e-05, + "loss": 0.1695, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.9449015259742737, + "learning_rate": 2e-05, + "loss": 0.1256, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.2726498544216156, + "learning_rate": 2e-05, + "loss": 0.0347, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.9403731226921082, + "learning_rate": 2e-05, + "loss": 0.0411, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.494799852371216, + "learning_rate": 2e-05, + "loss": 0.1758, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.06328664720058441, + "learning_rate": 2e-05, + "loss": 0.1062, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3632628917694092, + "learning_rate": 2e-05, + "loss": 0.1016, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.3595735728740692, + "learning_rate": 2e-05, + "loss": 0.0784, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 4.417634963989258, + "learning_rate": 2e-05, + "loss": 0.4862, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.02805049531161785, + "learning_rate": 2e-05, + "loss": 0.0466, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 4.830655574798584, + "learning_rate": 2e-05, + "loss": 0.1876, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5038827061653137, + "learning_rate": 2e-05, + "loss": 0.0328, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.22684521973133087, + "learning_rate": 2e-05, + "loss": 0.1935, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.3848518133163452, + "learning_rate": 2e-05, + "loss": 0.0681, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.2028435468673706, + "learning_rate": 2e-05, + "loss": 0.0263, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.9708964824676514, + "learning_rate": 2e-05, + "loss": 0.3365, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.17908445000648499, + "learning_rate": 2e-05, + "loss": 0.0079, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.852616310119629, + "learning_rate": 2e-05, + "loss": 0.3458, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.0022976398468018, + "learning_rate": 2e-05, + "loss": 0.115, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.0030856132507324, + "learning_rate": 2e-05, + "loss": 0.0937, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.739586591720581, + "learning_rate": 2e-05, + "loss": 0.0876, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.4155240058898926, + "learning_rate": 2e-05, + "loss": 0.1848, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.26595476269721985, + "learning_rate": 2e-05, + "loss": 0.3617, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.14563924074172974, + "learning_rate": 2e-05, + "loss": 0.0533, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8807713985443115, + "learning_rate": 2e-05, + "loss": 0.0801, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.007654379587620497, + "learning_rate": 2e-05, + "loss": 0.0284, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.7979322671890259, + "learning_rate": 2e-05, + "loss": 0.1033, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.034365035593509674, + "learning_rate": 2e-05, + "loss": 0.1249, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.558176577091217, + "learning_rate": 2e-05, + "loss": 0.0622, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 4.809686183929443, + "learning_rate": 2e-05, + "loss": 0.4131, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.05411836504936218, + "learning_rate": 2e-05, + "loss": 0.0124, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.2724485397338867, + "learning_rate": 2e-05, + "loss": 0.102, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.501843452453613, + "learning_rate": 2e-05, + "loss": 0.1306, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 4.625871181488037, + "learning_rate": 2e-05, + "loss": 0.2757, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.2811715304851532, + "learning_rate": 2e-05, + "loss": 0.0093, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 4.144275665283203, + "learning_rate": 2e-05, + "loss": 0.3001, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0693250447511673, + "learning_rate": 2e-05, + "loss": 0.1841, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.6654295325279236, + "learning_rate": 2e-05, + "loss": 0.0161, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.1108425185084343, + "learning_rate": 2e-05, + "loss": 0.0153, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.371753692626953, + "learning_rate": 2e-05, + "loss": 0.1497, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.03153133392334, + "learning_rate": 2e-05, + "loss": 0.133, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.105856418609619, + "learning_rate": 2e-05, + "loss": 0.3358, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 5.25473690032959, + "learning_rate": 2e-05, + "loss": 0.1006, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.019477348774671555, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.448582649230957, + "learning_rate": 2e-05, + "loss": 0.2571, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2207554989981696.0, + "train_loss": 0.13525307178497314, + "train_runtime": 81.0026, + "train_samples_per_second": 4.938, + "train_steps_per_second": 1.235 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2207554989981696.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..109a352a712bf803a6cb08909862a9cf2d2a9e93 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2574dbc2e58a8baf8b2ac1989b3c85d45614fa334a47c735da8473313139f7d6 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..1ed37eb35072b918e3c4ecc637da09f4ed14c178 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e1216b0119a17d4eafe0bda198ad6a4a4642fa312f06ea25a8c0a3fcec02339a +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e65135e66be0c9e5912967cda5764ee2ec722a6d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aa0885bdc195a68dd58c3603ff0bf771ef10ef803feb1dc4300ce6280ec029b2 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..54eb84da21876b8b2e5b18bfb0609661194a73dd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8ad18dcf0f371a361461049b662b1c195ddbbce08518b465aa17bb8b28909a12 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..ee226b2efd57be30953026d1277a339b6aae492d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:da8a3000a88b5d08b6da5d436233dbb486dbd813e4a3eded0f9e69ce08cd6e4c +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..479e7234255edcac8ffb32b888b89af5cc5f5f3b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c081ac6249ccbaa57376991da26c653644b5406c3086ebb8b6e9dc1199d8a5c3 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..91cebd25125bdd9189b3d77135ea336ad0d2eae2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c08cf70ea592067be0e1fb0e86622c284f277d0333118a64624defceeed7f77f +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f78ed75f10449092adf4da12b7da16db67cc5e71 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d7d51c2b1a241851ff59dd2c1ea86891630143b63f769a33efff6c6e04c244c3 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..308c4e3601f7c52b97177e660dd3636bcaf2e63d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.21480263769626617, + "learning_rate": 2e-05, + "loss": 0.0746, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.3009430766105652, + "learning_rate": 2e-05, + "loss": 0.0745, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.376713365316391, + "learning_rate": 2e-05, + "loss": 0.0378, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.5892430543899536, + "learning_rate": 2e-05, + "loss": 0.1315, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.222546577453613, + "learning_rate": 2e-05, + "loss": 0.38, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.29915037751197815, + "learning_rate": 2e-05, + "loss": 0.0455, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.9257538914680481, + "learning_rate": 2e-05, + "loss": 0.0672, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.5137848258018494, + "learning_rate": 2e-05, + "loss": 0.1237, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.3485954701900482, + "learning_rate": 2e-05, + "loss": 0.2787, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.09823121875524521, + "learning_rate": 2e-05, + "loss": 0.0118, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.8340469002723694, + "learning_rate": 2e-05, + "loss": 0.1119, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.9884753227233887, + "learning_rate": 2e-05, + "loss": 0.0498, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.620227575302124, + "learning_rate": 2e-05, + "loss": 0.1608, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.8484055995941162, + "learning_rate": 2e-05, + "loss": 0.116, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.313283681869507, + "learning_rate": 2e-05, + "loss": 0.261, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.32322588562965393, + "learning_rate": 2e-05, + "loss": 0.469, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.7812071442604065, + "learning_rate": 2e-05, + "loss": 0.2234, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.35307878255844116, + "learning_rate": 2e-05, + "loss": 0.1059, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.17220929265022278, + "learning_rate": 2e-05, + "loss": 0.0351, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.07107394188642502, + "learning_rate": 2e-05, + "loss": 0.0602, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.7156211733818054, + "learning_rate": 2e-05, + "loss": 0.0635, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.5215616226196289, + "learning_rate": 2e-05, + "loss": 0.4805, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4603126645088196, + "learning_rate": 2e-05, + "loss": 0.1289, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.2031881958246231, + "learning_rate": 2e-05, + "loss": 0.2488, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.2551450729370117, + "learning_rate": 2e-05, + "loss": 0.2525, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.6458418369293213, + "learning_rate": 2e-05, + "loss": 0.231, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.2966099679470062, + "learning_rate": 2e-05, + "loss": 0.1592, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.07342806458473206, + "learning_rate": 2e-05, + "loss": 0.0134, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.7996431589126587, + "learning_rate": 2e-05, + "loss": 0.1771, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.7205113172531128, + "learning_rate": 2e-05, + "loss": 0.0601, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.07053877413272858, + "learning_rate": 2e-05, + "loss": 0.1171, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.860989570617676, + "learning_rate": 2e-05, + "loss": 0.2538, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.0125670433044434, + "learning_rate": 2e-05, + "loss": 0.3254, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.05268147587776184, + "learning_rate": 2e-05, + "loss": 0.0847, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.3474907875061035, + "learning_rate": 2e-05, + "loss": 0.3749, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.5016542077064514, + "learning_rate": 2e-05, + "loss": 0.1097, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.22211211919784546, + "learning_rate": 2e-05, + "loss": 0.0272, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 4.742554664611816, + "learning_rate": 2e-05, + "loss": 0.5138, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.111788511276245, + "learning_rate": 2e-05, + "loss": 0.1762, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.1220197677612305, + "learning_rate": 2e-05, + "loss": 0.1376, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.5006320476531982, + "learning_rate": 2e-05, + "loss": 0.3421, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.09974780678749084, + "learning_rate": 2e-05, + "loss": 0.0239, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0017479404341429472, + "learning_rate": 2e-05, + "loss": 0.0333, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.10257686674594879, + "learning_rate": 2e-05, + "loss": 0.2601, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5315903425216675, + "learning_rate": 2e-05, + "loss": 0.0282, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.9760439395904541, + "learning_rate": 2e-05, + "loss": 0.0894, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.45380911231040955, + "learning_rate": 2e-05, + "loss": 0.0855, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.2826156616210938, + "learning_rate": 2e-05, + "loss": 0.1771, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.11086229234933853, + "learning_rate": 2e-05, + "loss": 0.0342, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.09485418349504471, + "learning_rate": 2e-05, + "loss": 0.2481, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5295755845697536.0, + "train_loss": 0.16151381969451906, + "train_runtime": 123.4096, + "train_samples_per_second": 3.241, + "train_steps_per_second": 0.81 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5295755845697536.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..fdea59386c90e8d1880c80d577d7353dcebce336 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:27e29a80d15bc2a08583231744993f832a478647dcda2b390ad0214ca2a0e629 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..3a1a4744cafee3cd9d625b2b55e5bc6278182430 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cc3591372429d3fbab8c99f085ad0e7003829648e69887676cabd42f9b97a46f +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..07699353a32799d35f5a992bbf66e1757adc1ff3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0c3f2a384039894d4570c58607f1f18f583d499b9db1e79441f94e7a979af6d7 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..f10d88045adde4d4e591b71e312358fd6951affd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7447b597d0b231b1425eecb052e17de43d986da912e1690b16ace291f8a7cb87 +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..ee264c7c34a77781d2996ee4da6004b4342b8613 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:32637e75a29315701481be568e6e8dca5ba79d652b2c5aa1e3083511ff01d5a4 +size 369837282 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..0b424a447a7ab5ea55ebb723d912c614593ce37a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:16ca3849b575643acbc9a02a05509e302e33ac7efc133d25ba22645b0d33631b +size 369838470 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2377937bd9a078425f74a71065ef2092b917d216 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:595a00545a1c9affb329399a54bc5a9f75ffdbf1216a79fcd752bcd12ddfe2c2 +size 369837282 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..6ff490428fe95c2b3107f0611809328eba94f001 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7ee506aa5127f3302c1600a4eaf814e0073c144ae80d5d080e6ae2259d44cda +size 369837282 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..e1a4babf6618fecbc1620e02a1833f6a53f3c17f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.05892818793654442, + "learning_rate": 2e-05, + "loss": 0.0175, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.0519106388092041, + "learning_rate": 2e-05, + "loss": 0.0269, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.16426512598991394, + "learning_rate": 2e-05, + "loss": 0.0042, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.14603987336158752, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.11088868975639343, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.4843077063560486, + "learning_rate": 2e-05, + "loss": 0.0424, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.01132002379745245, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.12457691878080368, + "learning_rate": 2e-05, + "loss": 0.0586, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.0747179388999939, + "learning_rate": 2e-05, + "loss": 0.0604, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.00530578987672925, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.02422078512609005, + "learning_rate": 2e-05, + "loss": 0.0198, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.10408667474985123, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.2423160076141357, + "learning_rate": 2e-05, + "loss": 0.0882, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.029950885102152824, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.08999351412057877, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.27807363867759705, + "learning_rate": 2e-05, + "loss": 0.0265, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.7234458923339844, + "learning_rate": 2e-05, + "loss": 0.0162, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.001890243380330503, + "learning_rate": 2e-05, + "loss": 0.0001, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.053295593708753586, + "learning_rate": 2e-05, + "loss": 0.0013, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.39068228006362915, + "learning_rate": 2e-05, + "loss": 0.0096, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.1167377308011055, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.6404411792755127, + "learning_rate": 2e-05, + "loss": 0.1432, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.010710666887462139, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.023454545065760612, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.0074742077849805355, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.03145710378885269, + "learning_rate": 2e-05, + "loss": 0.0401, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.680176019668579, + "learning_rate": 2e-05, + "loss": 0.0493, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.030208740383386612, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.0074619450606405735, + "learning_rate": 2e-05, + "loss": 0.05, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.014972697012126446, + "learning_rate": 2e-05, + "loss": 0.0131, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.27198657393455505, + "learning_rate": 2e-05, + "loss": 0.0036, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.014876780100166798, + "learning_rate": 2e-05, + "loss": 0.0063, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.03099854476749897, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.1861061453819275, + "learning_rate": 2e-05, + "loss": 0.0043, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.004303574096411467, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.2505335509777069, + "learning_rate": 2e-05, + "loss": 0.0043, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5974893569946289, + "learning_rate": 2e-05, + "loss": 0.0129, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.009959039278328419, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.9310294389724731, + "learning_rate": 2e-05, + "loss": 0.0204, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.005276673473417759, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.006976671516895294, + "learning_rate": 2e-05, + "loss": 0.0029, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0008760448545217514, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.020164260640740395, + "learning_rate": 2e-05, + "loss": 0.1046, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.009039030410349369, + "learning_rate": 2e-05, + "loss": 0.0009, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.2761610150337219, + "learning_rate": 2e-05, + "loss": 0.0066, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.5497651696205139, + "learning_rate": 2e-05, + "loss": 0.0047, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.00864244531840086, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.3139975070953369, + "learning_rate": 2e-05, + "loss": 0.0056, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.009200294502079487, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.025662733241915703, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2217994633609216.0, + "train_loss": 0.017491748332977296, + "train_runtime": 80.6362, + "train_samples_per_second": 4.961, + "train_steps_per_second": 1.24 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2217994633609216.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..45dadeb8d5bd14818e96d5921fc27dfa40e6574b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8529e94a1ae2be5e22e7c3df21af942df78fa99900106372bd99873f75f498b4 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..cfc72e5c790f0c49831b68f50aab5df42c3ef796 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e67119c49e82eb9904deb153909243779fff78f119f080b0f8f87f10f6eb91ad +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..43c9a69e7484d5e5bf3ddd23c1bfae5282a94021 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f73bafa62f6d4174bf37bf82cd9e0908ae704541413dbb1e0b0c164acfea49b5 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..a257bdedebf538fb8d52da5fff79f6b1d435f236 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2a10c9ddd9c07e423b7392eb989f1a46e35a8579f08a3ca64fb1667f78a53240 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..fa5ec54319892bf23e8c4a608477480aff40849a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8219133945bdea037992164bf71338b82beb507c3206b3331b426e7c5617d995 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..1ee5abfe3519e251b90b59772aef99f43b240ae7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:785dddd0e7cc682aa49929ba32eb67bcc0717d1e871dc8bbc54f94e3deb7a035 +size 794710050 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b6b603328b72f7d9d4a361e81011feafdf67f544 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6d4ec5bc49a921c7767b09e8ae3e548f2cde465ec5354cc786d5f1e1f505d12 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..6798423e2366f329d4f303138dfaa5318cec450a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5e2792ff72786407d72b0d543a6f32f8140eed5d3e30906f9d58e10c04c3e0e3 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..536e5c90ad167b6480ba73550dba0f6a9c08770a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.07659564167261124, + "learning_rate": 2e-05, + "loss": 0.0262, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.341485619544983, + "learning_rate": 2e-05, + "loss": 0.0753, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.4965511560440063, + "learning_rate": 2e-05, + "loss": 0.0793, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.40956762433052063, + "learning_rate": 2e-05, + "loss": 0.0218, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.453460216522217, + "learning_rate": 2e-05, + "loss": 0.4713, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.13105244934558868, + "learning_rate": 2e-05, + "loss": 0.0082, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.28581345081329346, + "learning_rate": 2e-05, + "loss": 0.0439, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.2174428701400757, + "learning_rate": 2e-05, + "loss": 0.0928, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.6102972030639648, + "learning_rate": 2e-05, + "loss": 0.1743, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.13349418342113495, + "learning_rate": 2e-05, + "loss": 0.1457, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.31393566727638245, + "learning_rate": 2e-05, + "loss": 0.0164, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.178712859749794, + "learning_rate": 2e-05, + "loss": 0.0159, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.8203631639480591, + "learning_rate": 2e-05, + "loss": 0.2669, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.1495300531387329, + "learning_rate": 2e-05, + "loss": 0.0146, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.6581099033355713, + "learning_rate": 2e-05, + "loss": 0.0253, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.5255067348480225, + "learning_rate": 2e-05, + "loss": 0.2814, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.0755738914012909, + "learning_rate": 2e-05, + "loss": 0.04, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3308757543563843, + "learning_rate": 2e-05, + "loss": 0.0341, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.22159229218959808, + "learning_rate": 2e-05, + "loss": 0.0682, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.38124316930770874, + "learning_rate": 2e-05, + "loss": 0.0566, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.2924318313598633, + "learning_rate": 2e-05, + "loss": 0.1485, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7537797689437866, + "learning_rate": 2e-05, + "loss": 0.0692, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.958396077156067, + "learning_rate": 2e-05, + "loss": 0.1181, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.5182420611381531, + "learning_rate": 2e-05, + "loss": 0.0387, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.097217082977295, + "learning_rate": 2e-05, + "loss": 0.1534, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.251325786113739, + "learning_rate": 2e-05, + "loss": 0.0118, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.45754072070121765, + "learning_rate": 2e-05, + "loss": 0.1314, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.04384413734078407, + "learning_rate": 2e-05, + "loss": 0.0353, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.01210832130163908, + "learning_rate": 2e-05, + "loss": 0.1609, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.15071886777877808, + "learning_rate": 2e-05, + "loss": 0.0607, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.6230262517929077, + "learning_rate": 2e-05, + "loss": 0.275, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.706148386001587, + "learning_rate": 2e-05, + "loss": 0.5074, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.920196056365967, + "learning_rate": 2e-05, + "loss": 0.7896, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.3980047702789307, + "learning_rate": 2e-05, + "loss": 0.1546, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.6564184427261353, + "learning_rate": 2e-05, + "loss": 0.0654, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.05167099088430405, + "learning_rate": 2e-05, + "loss": 0.024, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.11347205936908722, + "learning_rate": 2e-05, + "loss": 0.0146, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.327749729156494, + "learning_rate": 2e-05, + "loss": 0.6209, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.3568841516971588, + "learning_rate": 2e-05, + "loss": 0.0331, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.01498798094689846, + "learning_rate": 2e-05, + "loss": 0.0226, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.08554933965206146, + "learning_rate": 2e-05, + "loss": 0.0613, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.8002834916114807, + "learning_rate": 2e-05, + "loss": 0.1182, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.20100265741348267, + "learning_rate": 2e-05, + "loss": 0.0254, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.2844197750091553, + "learning_rate": 2e-05, + "loss": 0.1868, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5148859024047852, + "learning_rate": 2e-05, + "loss": 0.1406, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.514054298400879, + "learning_rate": 2e-05, + "loss": 0.324, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.5416432619094849, + "learning_rate": 2e-05, + "loss": 0.0403, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.053370654582977295, + "learning_rate": 2e-05, + "loss": 0.0035, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.309124231338501, + "learning_rate": 2e-05, + "loss": 0.1844, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.07195141166448593, + "learning_rate": 2e-05, + "loss": 0.0056, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5296221983866880.0, + "train_loss": 0.12968125820159912, + "train_runtime": 121.9705, + "train_samples_per_second": 3.279, + "train_steps_per_second": 0.82 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5296221983866880.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..d33af0612ecd9fea0d48b117a7911950d368f212 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1176e7eabc0014fc0e8daf21a64442ad3088f2e625190adc2ea416983a5d7666 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..f5621f076fdf494ef0b244fbd744ded26fdc4426 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b5ca3588854f0821dd4fba18b6fbac532ce4b662bee0157a99c2cf3fdd243dba +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..68825c0a7472f56d0668b1d4fec55134cb5f205b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d5119ec4c97a8af3d244574a870bca314a6febba90f9f4a60b53e114c48ec70 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0a70b5bec4216e44d93825fe2a8cf181cbaa59aa --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:baf7f3fc5bb74fad5820fef406e75cdc04d5edd36d6a6caad1dc488c2c477ef8 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..a9ec58bc95e6c480e69f4ced2ab750951f0972bc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8cd902e2d3b106a0fe67663e08d8f2b981c4fb29cb4764e4a3be6858bf6dcce +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..316607dde5822b1f30956aa8017f3f5e467f7e9c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7307ee3586b9d3602ad5cf6834f226d57fe9b0c728f3391dc0188f237f993571 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c11ffa076132885278e3da21c3641f69997c7f85 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:381de0fb80d2e36092422beb91bd5e3a4b1c4c468e3b3aac3cb910faa517b94d +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f16586e224f81b829158800e09520183d8103347 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1de3294beac426035cebcbc6df11bb0316933bf53fd1359534483b766a2f79eb +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..af7e3d3d0d1881702b3631d428e5d00886059ebb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.157696008682251, + "learning_rate": 2e-05, + "loss": 0.1332, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.4485087394714355, + "learning_rate": 2e-05, + "loss": 0.6074, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.118632197380066, + "learning_rate": 2e-05, + "loss": 0.4283, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.613885760307312, + "learning_rate": 2e-05, + "loss": 0.3208, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8376320600509644, + "learning_rate": 2e-05, + "loss": 0.1335, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.7430528402328491, + "learning_rate": 2e-05, + "loss": 0.5053, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.1812257170677185, + "learning_rate": 2e-05, + "loss": 0.1762, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.767461359500885, + "learning_rate": 2e-05, + "loss": 0.2423, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.003514289855957, + "learning_rate": 2e-05, + "loss": 0.2211, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.760804295539856, + "learning_rate": 2e-05, + "loss": 0.4779, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.33623385429382324, + "learning_rate": 2e-05, + "loss": 0.0718, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.3975290060043335, + "learning_rate": 2e-05, + "loss": 0.234, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.9884504079818726, + "learning_rate": 2e-05, + "loss": 0.5042, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.1897363662719727, + "learning_rate": 2e-05, + "loss": 0.1706, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.2567290663719177, + "learning_rate": 2e-05, + "loss": 0.0752, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.0995289087295532, + "learning_rate": 2e-05, + "loss": 0.1874, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.4498969316482544, + "learning_rate": 2e-05, + "loss": 0.3982, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.5457918047904968, + "learning_rate": 2e-05, + "loss": 0.0554, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.199607014656067, + "learning_rate": 2e-05, + "loss": 0.4482, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.6346216201782227, + "learning_rate": 2e-05, + "loss": 0.4236, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.8798087239265442, + "learning_rate": 2e-05, + "loss": 0.3401, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.3737882673740387, + "learning_rate": 2e-05, + "loss": 0.0695, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.7765726447105408, + "learning_rate": 2e-05, + "loss": 0.2063, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.32841381430625916, + "learning_rate": 2e-05, + "loss": 0.1702, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.8840925097465515, + "learning_rate": 2e-05, + "loss": 0.1212, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.376739025115967, + "learning_rate": 2e-05, + "loss": 0.2535, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.3070183992385864, + "learning_rate": 2e-05, + "loss": 0.2673, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.139783263206482, + "learning_rate": 2e-05, + "loss": 0.209, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.868559718132019, + "learning_rate": 2e-05, + "loss": 0.0716, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09497491270303726, + "learning_rate": 2e-05, + "loss": 0.1479, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.9386948943138123, + "learning_rate": 2e-05, + "loss": 0.3418, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.379090428352356, + "learning_rate": 2e-05, + "loss": 0.1964, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.6675780415534973, + "learning_rate": 2e-05, + "loss": 0.0843, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.8280346393585205, + "learning_rate": 2e-05, + "loss": 0.2407, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.3611595630645752, + "learning_rate": 2e-05, + "loss": 0.3873, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.42417266964912415, + "learning_rate": 2e-05, + "loss": 0.4575, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5183632969856262, + "learning_rate": 2e-05, + "loss": 0.0471, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.83855801820755, + "learning_rate": 2e-05, + "loss": 0.3059, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.6494832038879395, + "learning_rate": 2e-05, + "loss": 0.0684, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.3090503215789795, + "learning_rate": 2e-05, + "loss": 0.3991, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.9647715091705322, + "learning_rate": 2e-05, + "loss": 1.2744, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.1839742511510849, + "learning_rate": 2e-05, + "loss": 0.0186, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.3990212678909302, + "learning_rate": 2e-05, + "loss": 0.2927, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.228919506072998, + "learning_rate": 2e-05, + "loss": 0.3186, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.839862108230591, + "learning_rate": 2e-05, + "loss": 1.0483, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.1997346580028534, + "learning_rate": 2e-05, + "loss": 0.119, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.8532718420028687, + "learning_rate": 2e-05, + "loss": 0.2721, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.3296631872653961, + "learning_rate": 2e-05, + "loss": 0.0939, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.5096511840820312, + "learning_rate": 2e-05, + "loss": 0.6074, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.35642826557159424, + "learning_rate": 2e-05, + "loss": 0.0604, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5223464331902976.0, + "train_loss": 0.2861080741882324, + "train_runtime": 123.1865, + "train_samples_per_second": 3.247, + "train_steps_per_second": 0.812 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5223464331902976.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..98d55587008296052907daa78cd8b6a14a4a631e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:08ee40476ed8dc1300df7de5f4d4f9ed0ec5b33d354606a913706ee32b5154bd +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..828cae8c061c03f0720a755e03cd860b469f0309 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c80b14e81dc8159d09aa36bbb6ba5220572a0b993b55d7707b026b5c11145e1 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..1e39cb734efa0ed3db920e24563ebbbdd908e53d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e205fe3707eb4b7b99b213266f75da841ddefb2bb4207a9afea23fb3de6cfbe +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..319a8d5667831f3938afdead63d2e12823975777 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e5cf06eac1f7d90b3605a3c55956b2fb16e37157e2276486bf046955b28681e +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..fc413206e3e424841d09d035cdc3ff547a5b18a3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:230359ac164ba8c038b9dfaa06097a0d39c31216cf6c2cf24ba55d031e242710 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..45a931ea61f233a3912dbd81e55578af5f7bd306 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:11beba8777a3664c21409af2792a66dfad6a5b024b280583ae9128ac1bebeb2e +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c5ecb0dd36741ebefc0bf2a671a60ab6c3a28b66 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23d3221a42a42bb5425ada4f8f904d830f897115d668dde7d982164b535fe6e5 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d4c7f5d8361aa2cc4f03c4fd56f445b6179b3a8a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70b5ad907ac66284caaf95d47687060f7fb1e10878a87aa93f0948c2038a140f +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..ba0808340d9ff21bbf3c216cebc3d4c6b498d995 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.475977659225464, + "learning_rate": 2e-05, + "loss": 0.8276, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.6134572625160217, + "learning_rate": 2e-05, + "loss": 0.3521, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.2939074039459229, + "learning_rate": 2e-05, + "loss": 0.2956, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.369961977005005, + "learning_rate": 2e-05, + "loss": 0.6318, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.6359081268310547, + "learning_rate": 2e-05, + "loss": 0.4871, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.7912211418151855, + "learning_rate": 2e-05, + "loss": 0.7063, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.8888540267944336, + "learning_rate": 2e-05, + "loss": 0.7178, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.5904910564422607, + "learning_rate": 2e-05, + "loss": 0.5115, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.1876239776611328, + "learning_rate": 2e-05, + "loss": 0.1617, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.9772087335586548, + "learning_rate": 2e-05, + "loss": 0.3635, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.8389450311660767, + "learning_rate": 2e-05, + "loss": 0.2839, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.438183307647705, + "learning_rate": 2e-05, + "loss": 0.4785, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.3305557668209076, + "learning_rate": 2e-05, + "loss": 0.2179, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.377939224243164, + "learning_rate": 2e-05, + "loss": 0.3157, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.3438246250152588, + "learning_rate": 2e-05, + "loss": 0.4144, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.980118751525879, + "learning_rate": 2e-05, + "loss": 0.4166, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.5223909616470337, + "learning_rate": 2e-05, + "loss": 0.6768, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.84871768951416, + "learning_rate": 2e-05, + "loss": 0.5962, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.082582950592041, + "learning_rate": 2e-05, + "loss": 0.1653, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.5507313013076782, + "learning_rate": 2e-05, + "loss": 0.473, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.6427407264709473, + "learning_rate": 2e-05, + "loss": 0.5786, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.749948263168335, + "learning_rate": 2e-05, + "loss": 0.6438, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.9815380573272705, + "learning_rate": 2e-05, + "loss": 0.2357, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.6177049875259399, + "learning_rate": 2e-05, + "loss": 0.0677, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.6984091997146606, + "learning_rate": 2e-05, + "loss": 0.5889, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.3085687160491943, + "learning_rate": 2e-05, + "loss": 0.7312, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.08990608900785446, + "learning_rate": 2e-05, + "loss": 0.2078, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.170430064201355, + "learning_rate": 2e-05, + "loss": 0.2644, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.68681800365448, + "learning_rate": 2e-05, + "loss": 0.2988, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 3.003774881362915, + "learning_rate": 2e-05, + "loss": 0.5318, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.9555602073669434, + "learning_rate": 2e-05, + "loss": 0.3742, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.6162471771240234, + "learning_rate": 2e-05, + "loss": 0.2925, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.6824045181274414, + "learning_rate": 2e-05, + "loss": 0.4942, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.2239370346069336, + "learning_rate": 2e-05, + "loss": 0.8234, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.6490933895111084, + "learning_rate": 2e-05, + "loss": 0.395, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.2868715524673462, + "learning_rate": 2e-05, + "loss": 0.6958, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.4144070148468018, + "learning_rate": 2e-05, + "loss": 0.6652, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.8603956699371338, + "learning_rate": 2e-05, + "loss": 0.2262, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.5483289957046509, + "learning_rate": 2e-05, + "loss": 0.3428, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.382226586341858, + "learning_rate": 2e-05, + "loss": 0.1603, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.5020837187767029, + "learning_rate": 2e-05, + "loss": 0.134, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.9943599700927734, + "learning_rate": 2e-05, + "loss": 0.2414, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.215849757194519, + "learning_rate": 2e-05, + "loss": 0.5801, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0236562490463257, + "learning_rate": 2e-05, + "loss": 0.285, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5895534753799438, + "learning_rate": 2e-05, + "loss": 0.5038, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.0198700428009033, + "learning_rate": 2e-05, + "loss": 0.366, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.985944151878357, + "learning_rate": 2e-05, + "loss": 0.2727, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.0403239727020264, + "learning_rate": 2e-05, + "loss": 0.2075, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.098983645439148, + "learning_rate": 2e-05, + "loss": 0.3195, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.5579543709754944, + "learning_rate": 2e-05, + "loss": 0.2161, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5410845953622016.0, + "train_loss": 0.4167448616027832, + "train_runtime": 125.5759, + "train_samples_per_second": 3.185, + "train_steps_per_second": 0.796 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5410845953622016.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..7e4f9f3bffc4b8cb2452283d158bba61816c1936 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:797bc819bda19b311f831951fcb9b0a78063b0a94fc4c8637e35e42b2432056d +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..e79f2af12171e538ca5d8f0cfcf3b35427e68dc0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d6506c8992e964655ad2a63aa250455da7c3e1b52179821d9a318ce00bc05cfc +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e454dac2e0556618fb5398fc86c75aa174c497e5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:53334dc63b4f6e0e6181303bce1c2875c3c467ce2499e6eb105ae8488c280ddf +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..032a19b0a52e9fd11ea6c677849f8dbe0802cb82 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:13789a374c5f5d3195fb61082c831d68d249ff39013f2d88fce8ea6d504af4c5 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..20996d2db774d66c3467596aca674effa3b33f4a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b654abe2cdd29d7b5cc051f0088c3b03969469045a8e8907671c2f0b596f9998 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f46f182ce13aec1affc2c283da7126adfbc62ad --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76782afecd73adf6dde05e12ee5f57723e8e2fb4a62336c251f7d005a599e2e4 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..f787fa5eaa629df9d531156c3ba923d5236e1124 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1218cf071944dfb3673daa180c080259a45fd5b2f46576fb75a41748eaef312d +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..a5dbbf5085113d88ac373d1ec9855d3ee5b032fc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d53aa5627341962f86b85cdf073d09367c55c22cfc06ec801bf9f50a831cd26a +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..46fe367d571a37786bafa31e2f44953bd18914bb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.3529060482978821, + "learning_rate": 2e-05, + "loss": 0.1702, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.0109829902648926, + "learning_rate": 2e-05, + "loss": 0.3478, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.620620608329773, + "learning_rate": 2e-05, + "loss": 0.3848, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.0374565124511719, + "learning_rate": 2e-05, + "loss": 0.3742, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.7902807593345642, + "learning_rate": 2e-05, + "loss": 0.2034, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.244853138923645, + "learning_rate": 2e-05, + "loss": 0.1122, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.7450783252716064, + "learning_rate": 2e-05, + "loss": 0.298, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.866197347640991, + "learning_rate": 2e-05, + "loss": 0.5466, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.3544338941574097, + "learning_rate": 2e-05, + "loss": 0.1678, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.447248935699463, + "learning_rate": 2e-05, + "loss": 0.1831, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.7518043518066406, + "learning_rate": 2e-05, + "loss": 1.0105, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.8720446228981018, + "learning_rate": 2e-05, + "loss": 0.5106, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.3611409664154053, + "learning_rate": 2e-05, + "loss": 0.3339, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.362202525138855, + "learning_rate": 2e-05, + "loss": 0.453, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.9972494840621948, + "learning_rate": 2e-05, + "loss": 0.3614, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.4228821992874146, + "learning_rate": 2e-05, + "loss": 0.3484, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.6406632661819458, + "learning_rate": 2e-05, + "loss": 0.222, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.7479784488677979, + "learning_rate": 2e-05, + "loss": 0.4946, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.655400276184082, + "learning_rate": 2e-05, + "loss": 0.1937, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.6334545612335205, + "learning_rate": 2e-05, + "loss": 0.6426, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.8252699971199036, + "learning_rate": 2e-05, + "loss": 0.4136, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.9360519647598267, + "learning_rate": 2e-05, + "loss": 0.2771, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.7466654777526855, + "learning_rate": 2e-05, + "loss": 0.3679, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.851648211479187, + "learning_rate": 2e-05, + "loss": 0.345, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.0176005363464355, + "learning_rate": 2e-05, + "loss": 0.212, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.1303373575210571, + "learning_rate": 2e-05, + "loss": 0.6863, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.8946760892868042, + "learning_rate": 2e-05, + "loss": 0.2507, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5225472450256348, + "learning_rate": 2e-05, + "loss": 0.1405, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.280397653579712, + "learning_rate": 2e-05, + "loss": 0.3481, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.1729869842529297, + "learning_rate": 2e-05, + "loss": 0.4104, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7480144500732422, + "learning_rate": 2e-05, + "loss": 0.2377, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.8881944417953491, + "learning_rate": 2e-05, + "loss": 0.5049, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.4034955501556396, + "learning_rate": 2e-05, + "loss": 0.3645, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.1605172157287598, + "learning_rate": 2e-05, + "loss": 0.3645, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.1261088848114014, + "learning_rate": 2e-05, + "loss": 0.4014, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.1951546669006348, + "learning_rate": 2e-05, + "loss": 0.2464, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.501864492893219, + "learning_rate": 2e-05, + "loss": 0.288, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.0589303970336914, + "learning_rate": 2e-05, + "loss": 0.568, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.5424965620040894, + "learning_rate": 2e-05, + "loss": 0.3525, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.140094518661499, + "learning_rate": 2e-05, + "loss": 0.2811, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.7245244979858398, + "learning_rate": 2e-05, + "loss": 0.2252, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.445433259010315, + "learning_rate": 2e-05, + "loss": 0.3168, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.8986696004867554, + "learning_rate": 2e-05, + "loss": 0.1447, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.46963006258010864, + "learning_rate": 2e-05, + "loss": 0.2079, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.71208655834198, + "learning_rate": 2e-05, + "loss": 0.345, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.6946609020233154, + "learning_rate": 2e-05, + "loss": 0.2521, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.1732561588287354, + "learning_rate": 2e-05, + "loss": 0.4464, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.7984407544136047, + "learning_rate": 2e-05, + "loss": 0.3453, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.6913067102432251, + "learning_rate": 2e-05, + "loss": 0.0871, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.375337600708008, + "learning_rate": 2e-05, + "loss": 0.5294, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 6050325026832384.0, + "train_loss": 0.3463816833496094, + "train_runtime": 137.5727, + "train_samples_per_second": 2.908, + "train_steps_per_second": 0.727 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 6050325026832384.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f5a1c4db88acafdd78be2a4138abd2909da801ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c6939cdd349c05946d9ebbba539c0e807c08a9bfdfe869d583216b35d80e5ff6 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..da312095aa71a3cc81d23bb6e85beae2c9010914 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0878a25953b3f4270de127c2a5b5ae613dec1f22cedfd5e81119ffb6e3c35332 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..4c81cc398f2a1d3a59c0866501ba63b0b12b6279 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72772e2e48f9342632331f663dd3d09ed55bc0ac9c3e6412081ecf6cc09199c4 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..130805ae8d6e54ce743f7d8da2de126abea9bffd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:155f2261dae81732ac0d6dcdbb39a50b9743919bf3bd2d9844c337eb930cfb75 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..9a4635b8d7625770585ebbf28e8e7bef4c0d85b0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:01391b1cde98aeaa6e2f043f828312bd785edefba263a56734984509916e0010 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6c3a930b98f3db16f8b9a1d93ba575760313b2d5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b5071b3b81d35b0b94cbb0195d3331877a006925786113a30433a69bba57e46c +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..1480595a39b7519e931b904a2bd41fe9fe157f78 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5ccfc99d0491b73176d99d80808d63f845891e562e623f7a3727dc77d7142495 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f025d52f9758dad34187b76d253393757b223f03 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:458a8f809ae9db97ffec04ef6e236270aa3e8a54c648ac7cc41bde5980eda6c5 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..66376217e13e2e9363488f3734fa41eba148850c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.9547419548034668, + "learning_rate": 2e-05, + "loss": 0.1112, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.678523302078247, + "learning_rate": 2e-05, + "loss": 0.311, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.9979236125946045, + "learning_rate": 2e-05, + "loss": 0.1001, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.14306041598320007, + "learning_rate": 2e-05, + "loss": 0.0932, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.6064944863319397, + "learning_rate": 2e-05, + "loss": 0.118, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.11747434735298157, + "learning_rate": 2e-05, + "loss": 0.0359, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.004703618586063385, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.09892090409994125, + "learning_rate": 2e-05, + "loss": 0.0226, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.5486958026885986, + "learning_rate": 2e-05, + "loss": 0.0987, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.6591624021530151, + "learning_rate": 2e-05, + "loss": 0.1448, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.2570924758911133, + "learning_rate": 2e-05, + "loss": 0.1234, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.6353737115859985, + "learning_rate": 2e-05, + "loss": 0.1255, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.5532380938529968, + "learning_rate": 2e-05, + "loss": 0.1812, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.429884195327759, + "learning_rate": 2e-05, + "loss": 0.5669, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.07863155007362366, + "learning_rate": 2e-05, + "loss": 0.0349, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.7539189457893372, + "learning_rate": 2e-05, + "loss": 0.0505, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.4832432270050049, + "learning_rate": 2e-05, + "loss": 0.0981, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.21036870777606964, + "learning_rate": 2e-05, + "loss": 0.0365, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.7624675631523132, + "learning_rate": 2e-05, + "loss": 0.0334, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.9783471822738647, + "learning_rate": 2e-05, + "loss": 0.1502, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.1890951246023178, + "learning_rate": 2e-05, + "loss": 0.1975, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.68221116065979, + "learning_rate": 2e-05, + "loss": 0.2518, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.31784123182296753, + "learning_rate": 2e-05, + "loss": 0.1661, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.09300421923398972, + "learning_rate": 2e-05, + "loss": 0.9206, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.776052474975586, + "learning_rate": 2e-05, + "loss": 0.2554, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.780266284942627, + "learning_rate": 2e-05, + "loss": 0.1419, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.12001761049032211, + "learning_rate": 2e-05, + "loss": 0.0115, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.773501992225647, + "learning_rate": 2e-05, + "loss": 0.1068, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.19110868871212006, + "learning_rate": 2e-05, + "loss": 0.0167, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5428997874259949, + "learning_rate": 2e-05, + "loss": 0.0389, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.3614228665828705, + "learning_rate": 2e-05, + "loss": 0.0411, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.2800917625427246, + "learning_rate": 2e-05, + "loss": 0.6407, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.42213961482048035, + "learning_rate": 2e-05, + "loss": 0.0611, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.27532461285591125, + "learning_rate": 2e-05, + "loss": 0.0322, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.4998735189437866, + "learning_rate": 2e-05, + "loss": 0.2737, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.4233485162258148, + "learning_rate": 2e-05, + "loss": 0.2622, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.11715767532587051, + "learning_rate": 2e-05, + "loss": 0.1939, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.11692559719085693, + "learning_rate": 2e-05, + "loss": 0.0079, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.09763334691524506, + "learning_rate": 2e-05, + "loss": 0.5781, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.09278874099254608, + "learning_rate": 2e-05, + "loss": 0.0591, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.40508177876472473, + "learning_rate": 2e-05, + "loss": 0.0846, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.08843152225017548, + "learning_rate": 2e-05, + "loss": 0.0671, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.27156776189804077, + "learning_rate": 2e-05, + "loss": 0.0438, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.8901485800743103, + "learning_rate": 2e-05, + "loss": 0.0487, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.7410383224487305, + "learning_rate": 2e-05, + "loss": 0.0253, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.931152105331421, + "learning_rate": 2e-05, + "loss": 0.2241, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.4986700415611267, + "learning_rate": 2e-05, + "loss": 0.0263, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.226048707962036, + "learning_rate": 2e-05, + "loss": 0.2316, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.0296882763504982, + "learning_rate": 2e-05, + "loss": 0.2416, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 5.103370189666748, + "learning_rate": 2e-05, + "loss": 1.3198, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291433615425536.0, + "train_loss": 0.18076305389404296, + "train_runtime": 123.6772, + "train_samples_per_second": 3.234, + "train_steps_per_second": 0.809 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291433615425536.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..8872cee3d3526698ee178f20025b2fa862609779 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:28a809949e258f73948cd71493dff85b0b1c2cc310cbee96402f12a0b3f1b352 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..bbfef3cb6e8f2a77552c9a8b93d56f50dbdfa217 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7fff2057a5de3af28eedf971ddd8d6240cee3d36305f9ae11eae4705e0f13f0 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a9a078a02385e1d90bad01759c432031732f4c5b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8398604c76e92a39ebe48ed5323e0181da5e08d365dd1582d5c75117a683b635 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..53a02a260f1caa9efda3f9301bf43164c188b23a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e41dd2255ff51df9540722308a0d7c93830725b5806d862aba6b0fa47b1aea22 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..ee7e7b9f975f6c740064719a9b778d8def5fa7d3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8c34e2123dd74d2d08537019f09771a034ce91a259618ce935fe70f82bdbe58c +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d6d0323fd3c68dd357f5d325c143194ba3ab0c46 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e9331efe866e48b150509bb3452e21ecd7e1672bd6173b72f1023f64775b6054 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..bdc4b55090dbad804530e2fa205fb8ac5cc50152 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:25ce68ebcef38506edef805666836fad4e19c9fff89e1190fd9965d429bef605 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..85412fdabe98ccaaf9151e3c98bff53d34fe7b93 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d8c630b125cf7376a403a72f75d8255d0bdd7870ef3eabe2bdf2338bbef78b44 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..d855a2bdcacae4c5146cea70b0e2a416baf2e781 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.2891178131103516, + "learning_rate": 2e-05, + "loss": 0.4031, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.8194844722747803, + "learning_rate": 2e-05, + "loss": 0.3782, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.697237491607666, + "learning_rate": 2e-05, + "loss": 0.418, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.898367702960968, + "learning_rate": 2e-05, + "loss": 0.5092, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.76143479347229, + "learning_rate": 2e-05, + "loss": 0.5029, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.7311400175094604, + "learning_rate": 2e-05, + "loss": 0.0669, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.3763742446899414, + "learning_rate": 2e-05, + "loss": 0.3179, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.5790209770202637, + "learning_rate": 2e-05, + "loss": 0.5393, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.3445231914520264, + "learning_rate": 2e-05, + "loss": 0.8267, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.3094327449798584, + "learning_rate": 2e-05, + "loss": 0.3896, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.459745168685913, + "learning_rate": 2e-05, + "loss": 0.3904, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.5749486088752747, + "learning_rate": 2e-05, + "loss": 0.3033, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.521331787109375, + "learning_rate": 2e-05, + "loss": 0.3655, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.900700569152832, + "learning_rate": 2e-05, + "loss": 0.394, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.08201265335083, + "learning_rate": 2e-05, + "loss": 0.8523, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.601712703704834, + "learning_rate": 2e-05, + "loss": 0.3967, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.7393381595611572, + "learning_rate": 2e-05, + "loss": 0.3979, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.259629487991333, + "learning_rate": 2e-05, + "loss": 0.5354, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.1077476739883423, + "learning_rate": 2e-05, + "loss": 0.3627, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.0065797567367554, + "learning_rate": 2e-05, + "loss": 0.3203, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.3868894577026367, + "learning_rate": 2e-05, + "loss": 0.7838, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7546955943107605, + "learning_rate": 2e-05, + "loss": 0.2339, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.5606026649475098, + "learning_rate": 2e-05, + "loss": 0.4236, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.7913501262664795, + "learning_rate": 2e-05, + "loss": 0.3616, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.994629979133606, + "learning_rate": 2e-05, + "loss": 0.3932, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.1028581857681274, + "learning_rate": 2e-05, + "loss": 0.4111, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.9430327415466309, + "learning_rate": 2e-05, + "loss": 0.4153, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 3.508915662765503, + "learning_rate": 2e-05, + "loss": 0.8477, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.1863667964935303, + "learning_rate": 2e-05, + "loss": 0.4379, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.566255807876587, + "learning_rate": 2e-05, + "loss": 0.6694, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.6492022275924683, + "learning_rate": 2e-05, + "loss": 0.3672, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.4163440465927124, + "learning_rate": 2e-05, + "loss": 0.4521, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.1722304821014404, + "learning_rate": 2e-05, + "loss": 0.5698, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.7267134189605713, + "learning_rate": 2e-05, + "loss": 0.4579, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.75527286529541, + "learning_rate": 2e-05, + "loss": 0.6685, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.0502822399139404, + "learning_rate": 2e-05, + "loss": 0.2603, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.454706192016602, + "learning_rate": 2e-05, + "loss": 0.7451, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.317420482635498, + "learning_rate": 2e-05, + "loss": 0.4874, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.196631908416748, + "learning_rate": 2e-05, + "loss": 0.3491, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.947077989578247, + "learning_rate": 2e-05, + "loss": 0.2358, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.8898547887802124, + "learning_rate": 2e-05, + "loss": 0.4149, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.9015818238258362, + "learning_rate": 2e-05, + "loss": 0.4929, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.4217216968536377, + "learning_rate": 2e-05, + "loss": 0.2795, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.2374671697616577, + "learning_rate": 2e-05, + "loss": 0.3975, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.566948413848877, + "learning_rate": 2e-05, + "loss": 0.3241, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.6800459623336792, + "learning_rate": 2e-05, + "loss": 1.0186, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.2506022453308105, + "learning_rate": 2e-05, + "loss": 0.2722, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.9711779356002808, + "learning_rate": 2e-05, + "loss": 0.3019, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.0467472076416016, + "learning_rate": 2e-05, + "loss": 0.5664, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.102704048156738, + "learning_rate": 2e-05, + "loss": 0.9619, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 1.047142126845952e+16, + "train_loss": 0.4654170227050781, + "train_runtime": 179.1012, + "train_samples_per_second": 2.233, + "train_steps_per_second": 0.558 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 1.047142126845952e+16, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..37ee28956e3b27f2f16fe379187916d3ee7ce348 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:21f2714954363ff7315ab112b59af933f5601c42f621f8f9d9b2876d1858ee97 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ec4a54c862f646370577d893f3d7f55e2ee7bd4c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f9d1528072094ba36a57de9635fa67f9d9bc1fa7aae48e72974cf2a2a629c8cd +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..2ea3cf2072628be3a6b9e1b737c672b2af0c12db --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:388489e7bbe75dedc838e9ea4a2071fd46ba9b3a3373ab71153aad802eec49d5 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d64cb172a306b83e81f09e84c6bc12438f2deac3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9c378a5d551e43809b9a9bb5dc192c80e1ba4c7bb13c38ea63ea142640e35b0b +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..e71356e97ac0d8b76a205420fdd6f714d027f79d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aa49172cc1af86e508632cb0cbf640f8f41979444072e50b489b0fc75b66c2c5 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..28b5875c70fb4dbf64ce8872d7f34ebdc93038ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b4461418018ac8feaf10c3d9df1c13d68e5e1e156e995c968f78a2bdd17928d3 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..d14d8e029a1689b6554aff8c879f6975419ec3b2 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97cdec36046489ebec32ccba91b9603946152b7754607a410045b12e27c7d311 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..b47a01f3f937ed368564978e790dd79ae7eda374 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3cf0162cb0dd136035081d5c94d284763ae86dd78b8a5e49d0c7cccba9e76bdb +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..0a2ee8d3ab6e0b02a6a36fa46523f839cd384a45 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.687903881072998, + "learning_rate": 2e-05, + "loss": 0.0988, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.787522315979004, + "learning_rate": 2e-05, + "loss": 0.4229, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.792730689048767, + "learning_rate": 2e-05, + "loss": 0.1751, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.679871678352356, + "learning_rate": 2e-05, + "loss": 0.2049, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.9968682527542114, + "learning_rate": 2e-05, + "loss": 0.6607, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.11154723167419434, + "learning_rate": 2e-05, + "loss": 0.0462, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.30305716395378113, + "learning_rate": 2e-05, + "loss": 0.0806, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.373766541481018, + "learning_rate": 2e-05, + "loss": 0.3904, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.8321475982666016, + "learning_rate": 2e-05, + "loss": 0.1897, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.8013269901275635, + "learning_rate": 2e-05, + "loss": 0.3275, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.09388068318367004, + "learning_rate": 2e-05, + "loss": 0.2155, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.556795120239258, + "learning_rate": 2e-05, + "loss": 0.3253, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.7197142243385315, + "learning_rate": 2e-05, + "loss": 0.0395, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.1385817527770996, + "learning_rate": 2e-05, + "loss": 0.2727, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.3840588927268982, + "learning_rate": 2e-05, + "loss": 0.3058, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.6481873989105225, + "learning_rate": 2e-05, + "loss": 0.2737, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5026847720146179, + "learning_rate": 2e-05, + "loss": 0.0363, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.231173038482666, + "learning_rate": 2e-05, + "loss": 0.4164, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.6698359847068787, + "learning_rate": 2e-05, + "loss": 0.0594, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.5207974910736084, + "learning_rate": 2e-05, + "loss": 0.6881, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.40863677859306335, + "learning_rate": 2e-05, + "loss": 0.4891, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.5524911880493164, + "learning_rate": 2e-05, + "loss": 0.2413, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.3777424097061157, + "learning_rate": 2e-05, + "loss": 0.3219, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.1606188416481018, + "learning_rate": 2e-05, + "loss": 0.0088, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.371093213558197, + "learning_rate": 2e-05, + "loss": 0.4305, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.66108238697052, + "learning_rate": 2e-05, + "loss": 0.4266, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.038603998720645905, + "learning_rate": 2e-05, + "loss": 0.3341, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.46382036805152893, + "learning_rate": 2e-05, + "loss": 0.0356, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.1638120412826538, + "learning_rate": 2e-05, + "loss": 0.1022, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 6.998824596405029, + "learning_rate": 2e-05, + "loss": 1.173, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.4146518409252167, + "learning_rate": 2e-05, + "loss": 0.0616, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.5440212488174438, + "learning_rate": 2e-05, + "loss": 0.4675, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.4507945775985718, + "learning_rate": 2e-05, + "loss": 0.1424, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.585929811000824, + "learning_rate": 2e-05, + "loss": 0.126, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.0790464878082275, + "learning_rate": 2e-05, + "loss": 0.3941, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.328578233718872, + "learning_rate": 2e-05, + "loss": 0.3075, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.72244930267334, + "learning_rate": 2e-05, + "loss": 0.4807, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.40811148285865784, + "learning_rate": 2e-05, + "loss": 0.3495, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.214197039604187, + "learning_rate": 2e-05, + "loss": 0.3667, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.2521181106567383, + "learning_rate": 2e-05, + "loss": 0.807, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.4016233682632446, + "learning_rate": 2e-05, + "loss": 0.1562, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.2340495586395264, + "learning_rate": 2e-05, + "loss": 0.118, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.6554292440414429, + "learning_rate": 2e-05, + "loss": 0.3951, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.0585484504699707, + "learning_rate": 2e-05, + "loss": 0.3741, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.40072309970855713, + "learning_rate": 2e-05, + "loss": 0.2244, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.7712400555610657, + "learning_rate": 2e-05, + "loss": 0.159, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.8884232640266418, + "learning_rate": 2e-05, + "loss": 0.1118, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.2041546106338501, + "learning_rate": 2e-05, + "loss": 0.071, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.0212255716323853, + "learning_rate": 2e-05, + "loss": 0.1737, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.806027889251709, + "learning_rate": 2e-05, + "loss": 0.2726, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5502714666549248.0, + "train_loss": 0.2870355701446533, + "train_runtime": 132.9748, + "train_samples_per_second": 3.008, + "train_steps_per_second": 0.752 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5502714666549248.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..3ed872697f00a16155c4b02692011271d6456c91 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c07739b4a34838f92d60835244960fe514219e67bae12b90c13b0b5cb267f3b7 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..77a0d530f8e8b4bb68a0ad46d4b0c5a8db9e4a9a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b9986df492a0a5f86327ba5dc6a6c9dc7ffb42f8b6e0ecb530eb766c36d86bb8 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..6354e010ff1ca76c4d5026335b6fea43a21c47e5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:41703b8f5dcfb252a3dc07d7577d02e16553842c4b2e10a2f7a6788ef2ef3121 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..a73cc5ba2b47845fe04c33335113b885bd394c71 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3a8aae906e7d618c3e87cc8b6f5d1ff44174e2791571289b305fd77b8fcb7c39 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..b7f1a948e417f091b9b2d365c78897e63b26bd66 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1f0b0f0c1ba7d3574533677f466ea9e4956e692bd06ee81c6c5f6ca63808a5b0 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..76f257d6b2414102df398867f18334a96b8f632d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:884c5803bc070b547625d4166a4314581e80aaff52badd4a96cb89e2af093276 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b1865ea3c0e76571368d7994bfa1ffcb638017c3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d059b7b4d66a0b79c1aa2d98cfdd50991d407299a913af6907a631b9f93b6b11 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8af01d2d6e09c856111b03313f3c2515a7ab20ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7d81270ffd4d268fa435b9c2829d82512660774a47a5127aad52b36f80994d6d +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..4eee8db05ea11e545098dfc82180f80701a07f12 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.03161589801311493, + "learning_rate": 2e-05, + "loss": 0.0261, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.06569163501262665, + "learning_rate": 2e-05, + "loss": 0.076, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.1364307850599289, + "learning_rate": 2e-05, + "loss": 0.0451, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.851184844970703, + "learning_rate": 2e-05, + "loss": 0.2725, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.04096110910177231, + "learning_rate": 2e-05, + "loss": 0.0021, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.7173902988433838, + "learning_rate": 2e-05, + "loss": 0.0851, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.04092104360461235, + "learning_rate": 2e-05, + "loss": 0.1351, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 4.813397407531738, + "learning_rate": 2e-05, + "loss": 0.2361, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.0900509357452393, + "learning_rate": 2e-05, + "loss": 0.0991, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.4933342933654785, + "learning_rate": 2e-05, + "loss": 2.2706, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.4338135719299316, + "learning_rate": 2e-05, + "loss": 0.4478, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.4967397153377533, + "learning_rate": 2e-05, + "loss": 0.4336, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.861243724822998, + "learning_rate": 2e-05, + "loss": 0.1156, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.717741847038269, + "learning_rate": 2e-05, + "loss": 0.0449, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.1531623601913452, + "learning_rate": 2e-05, + "loss": 0.0496, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.04909582808613777, + "learning_rate": 2e-05, + "loss": 0.0067, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.8285107612609863, + "learning_rate": 2e-05, + "loss": 0.6875, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.5833683013916016, + "learning_rate": 2e-05, + "loss": 0.3143, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.6518778800964355, + "learning_rate": 2e-05, + "loss": 0.5566, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.0580483675003052, + "learning_rate": 2e-05, + "loss": 0.5301, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 5.321019649505615, + "learning_rate": 2e-05, + "loss": 0.4524, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.001976071624085307, + "learning_rate": 2e-05, + "loss": 0.059, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.9699909687042236, + "learning_rate": 2e-05, + "loss": 0.4021, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.0346460342407227, + "learning_rate": 2e-05, + "loss": 0.3232, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.2769052982330322, + "learning_rate": 2e-05, + "loss": 0.199, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.03202001005411148, + "learning_rate": 2e-05, + "loss": 0.1926, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.21090026199817657, + "learning_rate": 2e-05, + "loss": 0.2533, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5598428249359131, + "learning_rate": 2e-05, + "loss": 0.2701, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.8315578699111938, + "learning_rate": 2e-05, + "loss": 0.0572, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.06301199644804001, + "learning_rate": 2e-05, + "loss": 0.0091, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.967305064201355, + "learning_rate": 2e-05, + "loss": 0.1733, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.4686365127563477, + "learning_rate": 2e-05, + "loss": 0.5959, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.5326695442199707, + "learning_rate": 2e-05, + "loss": 0.0755, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.6539492011070251, + "learning_rate": 2e-05, + "loss": 0.0264, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.46062105894088745, + "learning_rate": 2e-05, + "loss": 0.0241, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.6404731273651123, + "learning_rate": 2e-05, + "loss": 0.2023, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.1135649681091309, + "learning_rate": 2e-05, + "loss": 0.0505, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.47032302618026733, + "learning_rate": 2e-05, + "loss": 0.0367, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.22391609847545624, + "learning_rate": 2e-05, + "loss": 0.1269, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.18165916204452515, + "learning_rate": 2e-05, + "loss": 0.0211, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.7258691787719727, + "learning_rate": 2e-05, + "loss": 0.5179, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.37550973892211914, + "learning_rate": 2e-05, + "loss": 0.0176, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.13239549100399017, + "learning_rate": 2e-05, + "loss": 0.0077, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.5001539587974548, + "learning_rate": 2e-05, + "loss": 0.0349, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.2907333374023438, + "learning_rate": 2e-05, + "loss": 0.1519, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.167206764221191, + "learning_rate": 2e-05, + "loss": 0.4848, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.21993699669837952, + "learning_rate": 2e-05, + "loss": 0.0076, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.9667516946792603, + "learning_rate": 2e-05, + "loss": 0.1463, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.16169828176498413, + "learning_rate": 2e-05, + "loss": 0.0089, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.1404234915971756, + "learning_rate": 2e-05, + "loss": 0.0088, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5329528591220736.0, + "train_loss": 0.22743694424629213, + "train_runtime": 124.7602, + "train_samples_per_second": 3.206, + "train_steps_per_second": 0.802 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5329528591220736.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..411db2ddedb664f9d54b389bb0cfe5ed3cea1b9a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2ea1c26dcd2620179c1696bd72ecd4b43131c9bb747d9b4bcdc75f9808301774 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..f062837cc2162e0bf7ee58ca9c57dc8245760339 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:42a436fc6fcf3ed567509586bfe21ed3a13d655aa6865024a4184a2858ac61e6 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..89248bc8b7dc33d65a09ee628291a08a9025cc9f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a2e3921dfd6e8bdf243c437e39a887c018e7f468a187cd621b8bd4b6a86c3448 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..784dd150f4901813183ff7113b8a27ce8968676b --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6b806798de128cb6a61c88b76894211de440db3e0cbc52d23285936897a9fe5 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5f689cd1a0fcb0c75645ea54372c544e2618f264 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ce78dca44ef984d66d2fc12170b656a5ed09571ced5f9c31141fab400c360326 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..ee741eb9aa9b13f3244f3abf74f18926154e3fd9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:278555ee8c706e193ac8a4df5157cdb82d1e94da336ce08adaaa1e70f9a79eb4 +size 794708086 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..a5adef1e2c779c67291621581aa3234d1e9636d4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:60536bae6e2bde58e5057174f1a2dd7bc36bfa51c3c32d71ccd30a963aadfcfd +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..7822b89645eec6fa7d10fee6ca0191442d9362d9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:006576943a0a78e251cbdc0f354f224b2ff9c4d01241fbbece044a3b5efac011 +size 794706058 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..fc160c69448adcb6f8f0e2d3b81439d7ce452997 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.2070704847574234, + "learning_rate": 2e-05, + "loss": 0.3326, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.49986547231674194, + "learning_rate": 2e-05, + "loss": 0.599, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.7320342063903809, + "learning_rate": 2e-05, + "loss": 0.0711, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.0681601762771606, + "learning_rate": 2e-05, + "loss": 0.116, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.14765986800193787, + "learning_rate": 2e-05, + "loss": 0.3116, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.4770420789718628, + "learning_rate": 2e-05, + "loss": 0.1411, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.0434364452958107, + "learning_rate": 2e-05, + "loss": 0.3877, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.335043430328369, + "learning_rate": 2e-05, + "loss": 0.5823, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.5816001892089844, + "learning_rate": 2e-05, + "loss": 0.1542, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.3459644317626953, + "learning_rate": 2e-05, + "loss": 0.2074, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.9057172536849976, + "learning_rate": 2e-05, + "loss": 0.7776, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.639756441116333, + "learning_rate": 2e-05, + "loss": 0.4679, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.879995346069336, + "learning_rate": 2e-05, + "loss": 0.4778, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.8578739166259766, + "learning_rate": 2e-05, + "loss": 0.0738, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.19263575971126556, + "learning_rate": 2e-05, + "loss": 0.2038, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.1069609671831131, + "learning_rate": 2e-05, + "loss": 0.0066, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.6667414903640747, + "learning_rate": 2e-05, + "loss": 0.5757, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.8021610379219055, + "learning_rate": 2e-05, + "loss": 0.1797, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2107805609703064, + "learning_rate": 2e-05, + "loss": 0.0582, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.631878137588501, + "learning_rate": 2e-05, + "loss": 0.3417, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.019948959350586, + "learning_rate": 2e-05, + "loss": 0.2786, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.1210811138153076, + "learning_rate": 2e-05, + "loss": 0.6068, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.6120047569274902, + "learning_rate": 2e-05, + "loss": 0.0555, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.058852195739746, + "learning_rate": 2e-05, + "loss": 0.4345, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.4894479513168335, + "learning_rate": 2e-05, + "loss": 0.0746, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.4978599548339844, + "learning_rate": 2e-05, + "loss": 0.0853, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.5401645302772522, + "learning_rate": 2e-05, + "loss": 0.483, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.13926281034946442, + "learning_rate": 2e-05, + "loss": 0.0135, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.8264352083206177, + "learning_rate": 2e-05, + "loss": 0.3462, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.6246852278709412, + "learning_rate": 2e-05, + "loss": 0.0841, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.042649984359741, + "learning_rate": 2e-05, + "loss": 0.1817, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.197062373161316, + "learning_rate": 2e-05, + "loss": 0.215, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.1000559329986572, + "learning_rate": 2e-05, + "loss": 0.1808, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.8150622248649597, + "learning_rate": 2e-05, + "loss": 0.0903, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.11400559544563293, + "learning_rate": 2e-05, + "loss": 0.0203, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 4.306662559509277, + "learning_rate": 2e-05, + "loss": 0.4461, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5178279876708984, + "learning_rate": 2e-05, + "loss": 0.1704, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.102329969406128, + "learning_rate": 2e-05, + "loss": 0.318, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.3291153013706207, + "learning_rate": 2e-05, + "loss": 0.1399, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.892282009124756, + "learning_rate": 2e-05, + "loss": 0.4497, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.8690320253372192, + "learning_rate": 2e-05, + "loss": 0.3181, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.9134912490844727, + "learning_rate": 2e-05, + "loss": 0.2393, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.1426920890808105, + "learning_rate": 2e-05, + "loss": 0.2238, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.4004942178726196, + "learning_rate": 2e-05, + "loss": 0.385, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.8965432047843933, + "learning_rate": 2e-05, + "loss": 0.2575, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.275634288787842, + "learning_rate": 2e-05, + "loss": 0.2717, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.5766851902008057, + "learning_rate": 2e-05, + "loss": 0.5264, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.31244802474975586, + "learning_rate": 2e-05, + "loss": 0.0487, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.6319400072097778, + "learning_rate": 2e-05, + "loss": 0.265, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.15896812081336975, + "learning_rate": 2e-05, + "loss": 0.0671, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293891360129024.0, + "train_loss": 0.26685415744781493, + "train_runtime": 131.3376, + "train_samples_per_second": 3.046, + "train_steps_per_second": 0.761 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293891360129024.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..3d29fdeeffdec57469bf5a603937d0d965b081c7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b3a88b2f88f775f842bf150c183ad9d9837519c38dbe164686ac4b8852a31202 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..71a1f1c2afba03484b09d97451454e399e3f67a7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1e9b1094397157af64a54a9c706f5f2d65e5dfbd531be23f0cc4613ec2a81fa8 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e346e613ecfd38e5608f06c2b4bd594fb610bcbf --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:87bae241e3bfd935a0ac4a6d024715205de26348769b2520dd4544c7cf03f704 +size 977003696 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1ae27704839e031ac54a058b2c7391ecbfa57244 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:65671ced59b408a3716f129cf6a43c36bd89796b608960897b50da416b158f75 +size 895433828 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7b4ee45cba23a2e197124c6e178f0428fb11f25c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:52cb3c54a5b3402d4638f58ba1cbbf1ee3c0e076efdbe1210e211702cfd07acd +size 977003696 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7d340ac41f3b8a2e938a7e5a200ba7fd3a33ae47 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:43a94aa41ca61590caaeac9fc4d35cb33b00d90a8f6b35ee8f105af9a960bac5 +size 895433828 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e72a04f6b718658dfb2e1c655d5c000c4640cbb6 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:52822fe3e1b45684133cc918d4f91527968383d83b4bcc62c1d2aa08ce365a82 +size 977003696 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..069919ee2d53e2ea6e97a9c2789eba274e181b4d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fc508c031f35988811392d1b04b14f92587fd34659933e5ed3a93d9f2455de12 +size 895433828 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c157ce9c14363c13e12d568dfa160ffdbad3671f --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:71623311758f31d3b370cb653c0dd281cf880efe7d8dfb673d953f77c57fc67c +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2906c37e583e1a64ca57443eba451c932b904187 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:42ff5fae8fd96b8b263e063404faaee2b2b9feb032b7daef65f5335434f4dec3 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d6d191083be43e47c05de25317b16cacbe54e6ef --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cf779d2bcd236d12af2e509199b73d7bb0a6a11b55a9d75e0eb8b3820d449a02 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..95f98551165c8725b5f117a9c149c420ec8f3397 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b1346ce46c0673457a22966e61c87b9b6b691fae5816e10e0ae4e4648c4c0740 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4332c676a41e08414356ef13958c45b874a3c073 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2d1c70219cf5364464aacc8dc48f22e17e89cce1cc2352e1b3b761073785a4e0 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1c9304a6b97586b3875114048e0bafc7005a8fa5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4d818af1f4ee75462b16d9a21cd430f0ae4143c90cb81743f754f924cf7d62ba +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..63b5d66d5b3277568d6b68c025e5443cfd4c4df5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20f2dc23b4efbca2d439918999aaacf9e5a72458ac9e9b2a8815dee55e6654ad +size 977003696 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c6eeda95c221de48a9092d0fe9bdb01ad7728601 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6a6468fe0d8bcfe50a937676d8505daf4ed488c80a3e1bffe1783176ff691d40 +size 895433828 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..dc252aba80837d26adc4d174dbe6f4ecf82759c3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9bb0404577347a4c9177fa7ff01c6f88d75cf3f88c9f0d345bf327dd682bb7bc +size 977003696 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..df14de4a6d47da4b4065113bb6f08b981501e6e1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9cc503ea4454164a1c282966564a626926d24864493708cdb77d5229e1026350 +size 895433828 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..baf5562dcc67e3e36c3f56f9216f786a0cb98dfa --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e33492cc7f79e2687f68657efb10e232182b8f24a48ae3e519830516fc6000f7 +size 977003696 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f0b988bb648e320a9f842c4d30710346934fc298 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d97fb4083fc939a86c5a368a24674be9d3ba72c120e92c2e0cf3322c1787bab +size 895433828 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..912cebf718bea3f100db739076ca663188f458fc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba10186ac8a3bf08e6085e91a9fc828c7277cf7459af4c099b451e02a9649057 +size 977003696 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..defba798ffbd8ef29349e40ed82e8ca53bc6d289 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f68a77d9c82781834bee297861ac03cb79ddde789edf35e562f9f1d301f8f198 +size 895433828 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..76798b86115e4e3e544f983ac464b816d9f4d220 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4dca8ddd7e21e88bc8dede08d1226c2241f0c21f1108db56d0c30dfac5e2b029 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2008010f53a6d2028294cf915856a16289be86d7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d305faab855c8756b48cdfb69ee109b287e8a553ec4ef8ded35c291dffcba6b0 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f8403c87b0a9b2d1b07360a07d0a4e31a215fe24 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e82c4e7f3b1e222f74e67297d34d9e32e8bea7d9510708f7fc340b571787f05 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e290155e369ac4d4c528ac50bc47f4bfa1dfd0ab --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:db0185a5446d342da4f1be9dd25e27a48d0ee4aa7506e843ea4a2aff6d6c37bb +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..42863f9eaa78a20e0b3921fbcfcc4a6307f96cfb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:444d1b0633f4d5ee9e3b16bf911548ed1ffb1359fad4ac49ab917ab74d9b17b6 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2c93f15357f6cb7c086fd154ac61dfb47cb27930 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:07044cbc5b5526c37ad9c1c2105cb6d69fd08a2ab836414b374729323783a1a8 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..67c0fe4685eb13c3b60c0084ea3bfa5bacdffdb0 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:995e5ff5c9d257b6c2f014817e43a0f9a328fb68bad1ff04b881b5e01d3b2467 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..433518c0e89d2518cf084d5d7ebbb9d6d7bdb033 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56214a2e30bf71117b832fd57ada26da51e9a12d417d7ad1edafc37114a13f8f +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a5cb0f45e17ea32d377253c95a3431e58d9d77dc --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b7253efc7c5347b9e1ebabf224439f5abccd29a0d904e50bd29b56090184a66 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6790a213df953b8e74dd5837820c27b8df7b68ef --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8444f87c1019cec4a9557c24434442e86764b3df8113eaa78f5e03c5af706e5 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b8c52ce259931f4e5c9aba5b1e06c9ce7208a2eb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:204752776727e51e74703078e329598562eee39f9303c5b45ded668e98fbd4db +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..066623c162877878493d31fd6cd41926e0440648 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:86f054b4a2de9aef6c2d7aa61fff0a99ea053d909bd09a427171b5b1796c041b +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..05885a12f70d8b26e73edbad8f28de9bfb0806a7 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:957a1ec0d5d18bb22afadea711338bb255c1ee45c5d1da917fe66df410d5f4bb +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..3b898cca3983af1fdb8a24cc866676d51fc57f4a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a112c5900fe00506308a79204c903aaa57d7cfc883f8dc076b1de7210b9a82e5 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9a9e6c6c7dd47ad365de78aaca7a3248607671bd --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7a483c3a620a307dd4dfbd59a7456c53bd5eb1605e0c20a2baaa44b7cb57d4d5 +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..17e2938791659c093e73f5dba7effccaf6c39c9a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:54cc55a7029ad94a3162910da6d8c36fc6b13d7a78b2dffac704e019d2574ec8 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7930b92e82569304121d53a521dba3abdd75afcb --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:96e6a687dfeb3477abc6763c4dad5a9bb876fe445bb9e339ae8c361ca2b4628e +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..fc9965915adbe09953fe707c693aa92b0d8817b8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a946e72949c9c4a8b3f6554288bbeb981dc3457fe177bce23c80fe05f194e118 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c53c2a0da9f1b66d716841e6f62a0649335aff33 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f60e2f4146170a71034c41a3c428fec63daa35b7aa6739fa36a2baa2786ae00f +size 2224279152 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..21c72659c60058506bd5fd4a130fd2c6b052a4f1 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf14fd7412d078e92d2acd44dba6c3e679b7747ab97050aa5e2546b1f786e087 +size 1159540164 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/latest b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..8480f3290fe02dcc8ad3f411dce4bb771607f136 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2a7a8d47edc80b205a0230dd0d312c16f3cc81aa0b01bd0d21718b6314735bb +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..120369764554a4d6cac58078d8512bf359607829 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf1dde35905b1ad125edf778aff1f9c3576ae0715c8c9876f665a63efcf18cd +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..64b454d945aa2a842c509c5cebe5ce0706d83a0c --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb23ed1efc2080025ddd7ca376ff771b3f058afa24e048f5cb3cf99201aa1e03 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ad8d07954c41361e73a0c9c7d3b4b46d894b188 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8caca9dd439d57dbc5054f65d33ebbd7ac694940b2b9dbe5e8cdf68881137a2 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f05da1a0577050cea961f470b13842ab5f1d3f9 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafca5706ac84f575d7648ee4786543561523e58cea8c01d2214b03d300bb988 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..351a4d2a284717aec90435f50067e1bb7002e9e4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:341a5a225c4c17a621d4f60caf4c34834000ea74e6f3c8b4f5aaada72ca1f547 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..08cfb19e58b68e297c605b87165d4acd241054e5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb378347913e2dd7be50fa6f9e3bd9e481f3897c2c9d31c597fe9164efc50bee +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..aeee08397c6123d33648f6703464379a4b7f85fe --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70c677ae195453774b8815440fee856ddb773c322681ac422ebc6e8abd5a58d7 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9935621fdc6da9e4abad379d97303b5ebf4f8f5 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:393c3d5a15b40b2ff049c254f5b09f25d226b0971827b4e7460ca58f9ff7b76f +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..2766cf1a4da1390df01ea4146b4a3a7a06e1434e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dc77db94ac5b4746eebf2c850a9425da9a425a71108790d6b7f5f2654b614a8 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..34bfc5121e29fc587c3f574c17d8006189d3021a --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47a94b88b812f21101345aae20a53cc10e0b346660c2cef669c77b201d06d508 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..4be93ba48f85d89c8cbd0e7426b79a10f51bc786 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20b9a2e67fdcb4051cebf52cba2ed0d47facb4dc69170f35a3483e3601bf25a4 +size 704650017 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..b42d717305b30cd186c1c408b8719020c0658a5d --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fd7458877767da0a36b80344cd6ec7eae8fbc8e11d07f8bfce3b4d5f5f239a9 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..01842d1316f21c3326c5735c6257b499da7ca458 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e537f056156082569ae0819abd06beb04b46581d0c63843bd68926e0ac2bff0e +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..3af9fe115499486166c3ef6d0a7bd27661ac08b4 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a9649d6997f84508e5f566393e91e42e03a31aca55878bacd7ffd2e134bbfc +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..abd1e557c23a8769870b0beedaf7df93836a101e --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7520bace4f501771d879621facd1b44f8db1ad7385a653e1535c7642465e67d +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..ebc846b7a5c587ced387cca85c4065e921c8cbe3 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1951a910d1b7587459db2110e8ac43cbc906a34267372eb21d2244a900eaec46 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..5eb1370d3fa3177cb4fb710882f4b7eb706817e8 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03f40a0df2437812b9e6de4e5d4e30014acb71344efa8132f19e1fc0f5925c82 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..385c599585a0973f240a027d28e985ca72d6b721 --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3ef6d8247bc37b017895c9e5bbaf32f2db500f94f66b4ebade5a4a5ba7a271 +size 704649992 diff --git a/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..219adb8927937a4597cd82efbf7ebb4c8a7519ad --- /dev/null +++ b/client_states_feddat_feddualMulti05pqfullfreeze_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:277fa5f0607d993e80153171e4760f999c4ee6a65452ea62694d0b9b03818f22 +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..2be2434e6cb878d4ce7ec6da95a82d966cf13d10 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:45aec727077880636924f444e2aee691a98f5fd67016fc2c3519b6eee4327e15 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..3bc0bf9d3aceb3cdb6de7c479a89417fe5f56b46 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d626fb335365f22d34ba565ed0138c03c207070af99209b5eba8f18180951a30 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..fb7dcb23f759e55d9a832b5a4977f1bf5d627e9b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a1810e071cd1a97024b5cc39e5cba2d576c09ec414d8f4c001cfbd37b84f0f2a +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..07e598b49a185241319635f3e5eafd83d9e3a071 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebdc88b3d2bf5279bc1924377a5ca1bd55dfecfbb22085f16c6fda4903066b74 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..6a64559621350e4dbe1cacaf44815ff036044b50 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3916135729ad4d513f6807a8b9611b900ecd9bfed3e6f6d3d5689da55ee0955a +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..f558a4f6a103f56891318b370a2de3ec195a92a0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e29fcf967352e19f95b2bf9306c68a02dcd2bfbc81cb62697b901c65fd5b7a6 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..41fc2be863407cf4f9e0bcc491efa72b8da877e2 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d0fe1c9e3bc442c4b7a2a87c6f394bf4125d451d2e1d9b74ded13866d34e3c9 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..57c044c16cb428e276d64fad22c4e27e64f367a8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3dbb8cc29efd09ce40ffd895d26f00d85d83b8fabe2cb0c7513566afbbfe3bf2 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..85c3f6db94e5730574c7e44612a9f28d800b443e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.148728847503662, + "learning_rate": 2e-05, + "loss": 0.1813, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.8716338872909546, + "learning_rate": 2e-05, + "loss": 0.159, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.3308429419994354, + "learning_rate": 2e-05, + "loss": 0.0311, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.24848555028438568, + "learning_rate": 2e-05, + "loss": 0.0471, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.8675153255462646, + "learning_rate": 2e-05, + "loss": 0.256, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.7386302947998047, + "learning_rate": 2e-05, + "loss": 0.2513, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.9215680360794067, + "learning_rate": 2e-05, + "loss": 0.0877, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.5844316482543945, + "learning_rate": 2e-05, + "loss": 0.4912, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 4.701166152954102, + "learning_rate": 2e-05, + "loss": 0.9702, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.304316282272339, + "learning_rate": 2e-05, + "loss": 0.3051, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.11880875378847122, + "learning_rate": 2e-05, + "loss": 0.0115, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.116666793823242, + "learning_rate": 2e-05, + "loss": 0.4189, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.23135802149772644, + "learning_rate": 2e-05, + "loss": 0.2124, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.2951427400112152, + "learning_rate": 2e-05, + "loss": 0.1828, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.3236154317855835, + "learning_rate": 2e-05, + "loss": 0.2651, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.12805584073066711, + "learning_rate": 2e-05, + "loss": 0.0212, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.3826881647109985, + "learning_rate": 2e-05, + "loss": 0.1791, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.47790002822876, + "learning_rate": 2e-05, + "loss": 0.338, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.5241738557815552, + "learning_rate": 2e-05, + "loss": 0.2139, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.0058271884918213, + "learning_rate": 2e-05, + "loss": 0.0939, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.21513792872428894, + "learning_rate": 2e-05, + "loss": 0.0217, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.6434217691421509, + "learning_rate": 2e-05, + "loss": 0.2942, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4751804769039154, + "learning_rate": 2e-05, + "loss": 0.0941, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.38501685857772827, + "learning_rate": 2e-05, + "loss": 0.1973, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.2681269645690918, + "learning_rate": 2e-05, + "loss": 0.0951, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.4966673851013184, + "learning_rate": 2e-05, + "loss": 0.0913, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.37936219573020935, + "learning_rate": 2e-05, + "loss": 0.252, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.7417842149734497, + "learning_rate": 2e-05, + "loss": 0.1255, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.596775770187378, + "learning_rate": 2e-05, + "loss": 0.5914, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.1073426604270935, + "learning_rate": 2e-05, + "loss": 0.0067, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.8306203484535217, + "learning_rate": 2e-05, + "loss": 0.1684, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.5724231004714966, + "learning_rate": 2e-05, + "loss": 0.126, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.3684564232826233, + "learning_rate": 2e-05, + "loss": 0.017, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.9377674460411072, + "learning_rate": 2e-05, + "loss": 0.0807, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.7186830043792725, + "learning_rate": 2e-05, + "loss": 0.2055, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.3584306240081787, + "learning_rate": 2e-05, + "loss": 0.0969, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.505014419555664, + "learning_rate": 2e-05, + "loss": 0.8641, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.902169942855835, + "learning_rate": 2e-05, + "loss": 0.4279, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.45388391613960266, + "learning_rate": 2e-05, + "loss": 0.249, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.5876932144165039, + "learning_rate": 2e-05, + "loss": 0.3042, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.677873969078064, + "learning_rate": 2e-05, + "loss": 0.0923, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.3424782752990723, + "learning_rate": 2e-05, + "loss": 0.2009, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.5234422087669373, + "learning_rate": 2e-05, + "loss": 0.0338, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.9441454410552979, + "learning_rate": 2e-05, + "loss": 0.1114, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3529950380325317, + "learning_rate": 2e-05, + "loss": 0.0871, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.817584216594696, + "learning_rate": 2e-05, + "loss": 0.2101, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.827960729598999, + "learning_rate": 2e-05, + "loss": 1.0267, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.1382263898849487, + "learning_rate": 2e-05, + "loss": 0.4014, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.435941219329834, + "learning_rate": 2e-05, + "loss": 0.2686, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.2399977743625641, + "learning_rate": 2e-05, + "loss": 0.0653, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291730211438592.0, + "train_loss": 0.23046264171600342, + "train_runtime": 129.6659, + "train_samples_per_second": 3.085, + "train_steps_per_second": 0.771 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291730211438592.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..d05a00d0234bdfb615b5352d43c2f801f7c437b6 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d2d1d9822d6b34a9379feda4e679396ef69fa9743c42b92ba143b660b6978777 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..bae8aa7e24dc838b1e7d5cdb0abf7e975f07ea67 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f573a8cb4d505db5c9231391d54764e544911294aba12083cc37d0b3358ba300 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..7c76422e8daf5bcbda7fd4650b6afa095de01ae3 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cc9358b39ef725af4db79728716a9eb56fce94cc8c88a883134f31a8cb40f091 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..dba7dbdb3a88aab92e34c847ffebeb81d0ac2092 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1846db53dbefb138cfcc3bd6bb612a11686b9f5bcadb7ab2fc8f7aa8c9dee7ba +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..209cfc2c2d6a1c660dbad6947c1cad032d85ca39 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d17d716d6eab63eebf9d710454840a636869540a460b5b7a55d1337f25a3200 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..73678005c9436a1c872eae6b3174305046c3f607 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c528fad358824c3b465c26f0a5a4bd20e0acf090388075da365cf4c7754f8d85 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..6146c45804e2c66d30ea374da592c2eee238ad05 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:603499abe07414dc4f32f2860187ac4533fb858a31c538e75de5935fa42c955c +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..3b904347406663aaef84e02aabc521446a9c10f9 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:445479038aaeb25fd3e87a635a285d7f9d5009100172a5fec11c08b3ee5e24f6 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3f23af232541fb5ace467f030a0ab68ac98d5ebb --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 5.796288967132568, + "learning_rate": 2e-05, + "loss": 0.3016, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 7.015511512756348, + "learning_rate": 2e-05, + "loss": 0.2218, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 4.8412861824035645, + "learning_rate": 2e-05, + "loss": 0.1755, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 4.48398494720459, + "learning_rate": 2e-05, + "loss": 0.638, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.0685272216796875, + "learning_rate": 2e-05, + "loss": 0.2795, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.4700462520122528, + "learning_rate": 2e-05, + "loss": 0.0172, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 9.151591300964355, + "learning_rate": 2e-05, + "loss": 0.4025, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.593291759490967, + "learning_rate": 2e-05, + "loss": 0.1693, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 5.341460704803467, + "learning_rate": 2e-05, + "loss": 0.1168, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.699438571929932, + "learning_rate": 2e-05, + "loss": 0.2656, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 5.9027838706970215, + "learning_rate": 2e-05, + "loss": 0.2323, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 4.126018524169922, + "learning_rate": 2e-05, + "loss": 0.2785, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 6.42581033706665, + "learning_rate": 2e-05, + "loss": 0.548, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 22.329734802246094, + "learning_rate": 2e-05, + "loss": 0.7495, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 15.705430030822754, + "learning_rate": 2e-05, + "loss": 0.7097, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 4.727340221405029, + "learning_rate": 2e-05, + "loss": 0.0583, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 11.205009460449219, + "learning_rate": 2e-05, + "loss": 0.233, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3244923949241638, + "learning_rate": 2e-05, + "loss": 0.2953, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.8708255290985107, + "learning_rate": 2e-05, + "loss": 0.1222, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.382613658905029, + "learning_rate": 2e-05, + "loss": 0.3605, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.5260701179504395, + "learning_rate": 2e-05, + "loss": 0.1158, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.34585583209991455, + "learning_rate": 2e-05, + "loss": 0.6286, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.116807222366333, + "learning_rate": 2e-05, + "loss": 0.0779, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.69175124168396, + "learning_rate": 2e-05, + "loss": 0.256, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 3.3057827949523926, + "learning_rate": 2e-05, + "loss": 0.1259, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.87550950050354, + "learning_rate": 2e-05, + "loss": 0.8577, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 13.339618682861328, + "learning_rate": 2e-05, + "loss": 0.8965, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.8989115357398987, + "learning_rate": 2e-05, + "loss": 0.0364, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 17.739574432373047, + "learning_rate": 2e-05, + "loss": 0.1797, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 10.00617504119873, + "learning_rate": 2e-05, + "loss": 0.41, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 7.7322797775268555, + "learning_rate": 2e-05, + "loss": 0.4185, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 5.8162360191345215, + "learning_rate": 2e-05, + "loss": 0.1943, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.746706962585449, + "learning_rate": 2e-05, + "loss": 0.8948, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.803368091583252, + "learning_rate": 2e-05, + "loss": 0.347, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.91923189163208, + "learning_rate": 2e-05, + "loss": 0.0512, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 21.376502990722656, + "learning_rate": 2e-05, + "loss": 0.8813, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 7.3373703956604, + "learning_rate": 2e-05, + "loss": 0.2935, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 13.966336250305176, + "learning_rate": 2e-05, + "loss": 0.3101, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 7.807223796844482, + "learning_rate": 2e-05, + "loss": 0.7515, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.5656020641326904, + "learning_rate": 2e-05, + "loss": 0.6284, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 9.010897636413574, + "learning_rate": 2e-05, + "loss": 0.3883, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.1942838430404663, + "learning_rate": 2e-05, + "loss": 0.0101, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 8.877424240112305, + "learning_rate": 2e-05, + "loss": 0.4397, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 7.212399005889893, + "learning_rate": 2e-05, + "loss": 0.4805, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.383707880973816, + "learning_rate": 2e-05, + "loss": 0.1613, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 7.696774005889893, + "learning_rate": 2e-05, + "loss": 0.4431, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.709601879119873, + "learning_rate": 2e-05, + "loss": 0.6861, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 4.7617669105529785, + "learning_rate": 2e-05, + "loss": 0.1387, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 5.173598289489746, + "learning_rate": 2e-05, + "loss": 0.3298, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.6681382656097412, + "learning_rate": 2e-05, + "loss": 0.458, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2221568423886848.0, + "train_loss": 0.3613085556030273, + "train_runtime": 86.1162, + "train_samples_per_second": 4.645, + "train_steps_per_second": 1.161 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2221568423886848.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..ed8e736701ba7285f0b4f6bfee6a99d9a3da57f3 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f15f76030b28146dc1e49172d6dcc330aebab077fc562b98afc7df61d690587d +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..99907dec2dfe2187ef420ec9ef8f5d22e355416e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d601f5043c250bb5af2dee09710f6361b4809eb5c831750cf07bbe5dff7c4e2 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..1e94d3b535708fbf6ff646e549e109f76c8df10a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:343b2b146adeb5a2013c3893b4c3cd1cb148d454483d609e8521f166b9ffced6 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..cff8318833bf5cf22e5b5152fdbc699cd95504b5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d49dfcb45d65a82f680a2c8f42b59c53a77299f69a5811d79442b2b4e24502d6 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..9974649ba454c7577516aba1ed22c69e59658dc9 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8849318543c1b7ef8a5f04c875a70391a0898c76791f7a4793b8d602e61e76a8 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b1ee9fa5c2edaf1e44dd35f1606f72703a42bc22 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b71826f8081ee0f020d7e10654873a23037d7f849593914ca44a7bb65d2b1d3 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..9b3d55b14703574e43a0b426cc304caff22c27ef --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:268a761764f8edf4a59af030a6288fb70f7993a62e54eb2c34327f63bf5c2f53 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..c5e07d4e2a8513f3829799b5a16b6706b653cc60 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6669f67462ff663d28c819aa289ce92534bd5e320ff9ef1d1de5455b2e105e0b +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..a34ffe8c39092c65a1e1b9eae1a58220b0201616 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.601040840148926, + "learning_rate": 2e-05, + "loss": 0.543, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.710258483886719, + "learning_rate": 2e-05, + "loss": 0.4556, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.8245387077331543, + "learning_rate": 2e-05, + "loss": 0.3268, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.665682315826416, + "learning_rate": 2e-05, + "loss": 0.3752, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.678458213806152, + "learning_rate": 2e-05, + "loss": 0.7207, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.6499297618865967, + "learning_rate": 2e-05, + "loss": 0.4054, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 5.280274391174316, + "learning_rate": 2e-05, + "loss": 0.9185, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 6.322684288024902, + "learning_rate": 2e-05, + "loss": 0.7793, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.2779006958007812, + "learning_rate": 2e-05, + "loss": 0.5652, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.398438811302185, + "learning_rate": 2e-05, + "loss": 0.3821, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.3832911252975464, + "learning_rate": 2e-05, + "loss": 0.2917, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.0447144508361816, + "learning_rate": 2e-05, + "loss": 0.4893, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.464441180229187, + "learning_rate": 2e-05, + "loss": 0.3925, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.5187478065490723, + "learning_rate": 2e-05, + "loss": 0.3743, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.437949776649475, + "learning_rate": 2e-05, + "loss": 0.4609, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.3522839546203613, + "learning_rate": 2e-05, + "loss": 0.3853, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 4.390581130981445, + "learning_rate": 2e-05, + "loss": 0.5049, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.671212673187256, + "learning_rate": 2e-05, + "loss": 0.459, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.0405789613723755, + "learning_rate": 2e-05, + "loss": 0.5383, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.9937263131141663, + "learning_rate": 2e-05, + "loss": 0.4317, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.331000804901123, + "learning_rate": 2e-05, + "loss": 0.3372, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.584084987640381, + "learning_rate": 2e-05, + "loss": 0.5641, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.2591934204101562, + "learning_rate": 2e-05, + "loss": 0.4963, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.6916806101799011, + "learning_rate": 2e-05, + "loss": 0.2577, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 6.473905563354492, + "learning_rate": 2e-05, + "loss": 0.5661, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.360517501831055, + "learning_rate": 2e-05, + "loss": 0.4839, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.7887448072433472, + "learning_rate": 2e-05, + "loss": 0.3221, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.4815196990966797, + "learning_rate": 2e-05, + "loss": 0.2632, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 4.027200222015381, + "learning_rate": 2e-05, + "loss": 0.2756, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.1629983186721802, + "learning_rate": 2e-05, + "loss": 0.3251, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 5.690079689025879, + "learning_rate": 2e-05, + "loss": 0.4515, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.8319354057312012, + "learning_rate": 2e-05, + "loss": 0.3633, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.708216667175293, + "learning_rate": 2e-05, + "loss": 0.6401, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.19010047614574432, + "learning_rate": 2e-05, + "loss": 0.1907, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.5467565059661865, + "learning_rate": 2e-05, + "loss": 0.3305, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 3.2586264610290527, + "learning_rate": 2e-05, + "loss": 0.4099, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.2688286304473877, + "learning_rate": 2e-05, + "loss": 0.6831, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.742818832397461, + "learning_rate": 2e-05, + "loss": 0.3763, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 7.051976680755615, + "learning_rate": 2e-05, + "loss": 0.5883, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 4.374168872833252, + "learning_rate": 2e-05, + "loss": 0.6318, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 6.1318488121032715, + "learning_rate": 2e-05, + "loss": 0.7793, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 4.306075572967529, + "learning_rate": 2e-05, + "loss": 0.6494, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.338111400604248, + "learning_rate": 2e-05, + "loss": 0.3308, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.155094861984253, + "learning_rate": 2e-05, + "loss": 0.4678, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.02273320034146309, + "learning_rate": 2e-05, + "loss": 0.4815, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.6005859375, + "learning_rate": 2e-05, + "loss": 0.344, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.7294667959213257, + "learning_rate": 2e-05, + "loss": 0.4707, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.649329662322998, + "learning_rate": 2e-05, + "loss": 0.4263, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.5623779296875, + "learning_rate": 2e-05, + "loss": 0.4023, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.473285436630249, + "learning_rate": 2e-05, + "loss": 0.4419, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2192326831112192.0, + "train_loss": 0.46300411224365234, + "train_runtime": 89.1399, + "train_samples_per_second": 4.487, + "train_steps_per_second": 1.122 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2192326831112192.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5b41ca7225b773fad8c908a13b27cf8e7518b18d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:613109a85d4e91f87ffaa4720b1c4cee21aad747ead59d027492ddde734e2c8c +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..fb4f53dfb9ca79130c92c55223c5d5abdd69db70 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd316d2e5c69ed50ec88a2753dc3cd91e7a27fe327e3bfa5bc77409eef99915b +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..46aa1ad4a56ff3d40b75e5ef06b7c230104ee00d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0b1a7ed00a6aa6e0d454c3298eef6489b23061e8d0230fe76d57eb0d20c889ee +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..4b053d446d0b59659d922ced1a67624135219913 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:efb44e908e35c233fb8b4f8f355242b806ebbc561a275ab0b27b487facca6f15 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..72ce7cf71adaa96a529b18cfbe63a1443dc22e66 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6006120da3a32c4fa0938c921cd1a149270a8ff0f869d104464940ba3022608a +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d6b7b4ba0eaa1773c58f227468715a92f041468b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:752cb46a7679ff505cceaed1a9d361599cc042983f20b32c07cbab386baf7a7c +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..49bfc7613a6157554ca416c970eb84a98fcf3fed --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:75cc3d302ec847aa00562a85d922271d0658b60d2e437f4851d2ac3b80a6e9c8 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..619b1fdf1dc6d8d0b21bfef71be7099325419ee8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:17c3e16d52e7985452358d7e2ba2a79ce9f708daff7e3d05806196f5c51d7a3b +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..af5fe9bc41ae0661f0bad8c22435a3d50d803ca3 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.9585040807723999, + "learning_rate": 2e-05, + "loss": 0.0406, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.0025479288306087255, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.06834162026643753, + "learning_rate": 2e-05, + "loss": 0.0113, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.6298388838768005, + "learning_rate": 2e-05, + "loss": 0.0276, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.2625581920146942, + "learning_rate": 2e-05, + "loss": 0.0158, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.05499662831425667, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.448557138442993, + "learning_rate": 2e-05, + "loss": 0.6259, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.4906578063964844, + "learning_rate": 2e-05, + "loss": 0.2857, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.28847336769104004, + "learning_rate": 2e-05, + "loss": 0.1339, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.094550132751465, + "learning_rate": 2e-05, + "loss": 0.133, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.07939810305833817, + "learning_rate": 2e-05, + "loss": 0.1189, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.2217295914888382, + "learning_rate": 2e-05, + "loss": 0.0164, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.30864402651786804, + "learning_rate": 2e-05, + "loss": 0.039, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.7758409976959229, + "learning_rate": 2e-05, + "loss": 0.1356, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.030043045058846474, + "learning_rate": 2e-05, + "loss": 0.0125, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.06484036147594452, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.08921733498573303, + "learning_rate": 2e-05, + "loss": 0.0313, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.13179801404476166, + "learning_rate": 2e-05, + "loss": 0.0553, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.12339866906404495, + "learning_rate": 2e-05, + "loss": 0.0237, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.5287211537361145, + "learning_rate": 2e-05, + "loss": 0.0468, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.752568244934082, + "learning_rate": 2e-05, + "loss": 0.3224, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.029347974807024002, + "learning_rate": 2e-05, + "loss": 0.025, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.6244646310806274, + "learning_rate": 2e-05, + "loss": 0.1904, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.5810307264328003, + "learning_rate": 2e-05, + "loss": 0.0265, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.05023611709475517, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.06780068576335907, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.052989713847637177, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.18092650175094604, + "learning_rate": 2e-05, + "loss": 0.0064, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.017643742263317108, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.049797169864177704, + "learning_rate": 2e-05, + "loss": 0.0036, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.011705023236572742, + "learning_rate": 2e-05, + "loss": 0.0073, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.011477609165012836, + "learning_rate": 2e-05, + "loss": 0.0052, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.17514748871326447, + "learning_rate": 2e-05, + "loss": 0.0051, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7718544602394104, + "learning_rate": 2e-05, + "loss": 0.0332, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.007084485609084368, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.025074122473597527, + "learning_rate": 2e-05, + "loss": 0.2736, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.04243431240320206, + "learning_rate": 2e-05, + "loss": 0.0026, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.010357972234487534, + "learning_rate": 2e-05, + "loss": 0.0204, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.30428266525268555, + "learning_rate": 2e-05, + "loss": 0.1524, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.005521278828382492, + "learning_rate": 2e-05, + "loss": 0.0085, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.04301392287015915, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.06476859003305435, + "learning_rate": 2e-05, + "loss": 0.005, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0040264250710606575, + "learning_rate": 2e-05, + "loss": 0.0087, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.8152210712432861, + "learning_rate": 2e-05, + "loss": 0.1812, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.004585479851812124, + "learning_rate": 2e-05, + "loss": 0.0142, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.47056809067726135, + "learning_rate": 2e-05, + "loss": 0.0909, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.4582786560058594, + "learning_rate": 2e-05, + "loss": 0.4731, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.035407405346632004, + "learning_rate": 2e-05, + "loss": 0.4452, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.046511851251125336, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.0899503231048584, + "learning_rate": 2e-05, + "loss": 0.0321, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5287535144075264.0, + "train_loss": 0.0823326814174652, + "train_runtime": 132.1857, + "train_samples_per_second": 3.026, + "train_steps_per_second": 0.757 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5287535144075264.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..d55b54240dca5d4c4be7bbc030877ec947996ec5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8f21a43114310f1bdfdc5072b0b081541d29ebfa0af0c73d982d1d5b963512c +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..4b3569377faea4ca58ccbb64b0ce3cfa81c6ac6e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:85333e008393d3da4da8551f0b03e9b8434e7d3c005d58ec95b1e064fcdc4573 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..0c0b46267099058651ac87b04def34b954da2503 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:310db6ecf10ada5bc6af597e790683106999e87583356f931b3964d307636b50 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..92fff64ba29566916e66973f7d7d95fe41c43c57 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:67d26a38a898aa05e32f6bc9c016eb27831f169d87c1277df967304c1bc7502f +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..42a5d97446d9c9de0ce5e0bd8805337799ba5566 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4c803f65928bdaa26729746fc5a12526d33856e7cf075d6bb53f85454f1f9a20 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..9f03b4e5f86c204a6df0777f3101d24c7a8abf98 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4ef7ca9da10447ec27541bc01710d6451c51990fcb4b92e95888a625217c4bc0 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..eeb2a47e08464989370282bb4a1f11c0b4d40885 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4158246181c4a73fe31e26851c34cc35349b2c4c08dc6f0487a502f9c6440c8f +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..47f8fa7aed1e82f1e77983ca1a3b884ea2be328b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:31e9fac6ed8d872c19ed2ab24871fd24ab80c9df0c282dfa54a7560d8fea9600 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..dac2fc0ef9933216d11f52dc9981b4c8098643d9 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.079291343688965, + "learning_rate": 2e-05, + "loss": 0.2489, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.6754429340362549, + "learning_rate": 2e-05, + "loss": 0.2092, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.9001460075378418, + "learning_rate": 2e-05, + "loss": 0.4328, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.9351787567138672, + "learning_rate": 2e-05, + "loss": 0.0837, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.5929735898971558, + "learning_rate": 2e-05, + "loss": 0.0856, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.24269598722457886, + "learning_rate": 2e-05, + "loss": 0.036, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.664595127105713, + "learning_rate": 2e-05, + "loss": 0.0816, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.5634896755218506, + "learning_rate": 2e-05, + "loss": 0.2966, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.6558717489242554, + "learning_rate": 2e-05, + "loss": 0.2201, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.111650824546814, + "learning_rate": 2e-05, + "loss": 1.2049, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.995584487915039, + "learning_rate": 2e-05, + "loss": 0.5327, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 6.254846572875977, + "learning_rate": 2e-05, + "loss": 0.2612, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.1186509132385254, + "learning_rate": 2e-05, + "loss": 0.9696, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.5527901649475098, + "learning_rate": 2e-05, + "loss": 0.1653, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.8733394145965576, + "learning_rate": 2e-05, + "loss": 0.2669, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.123582124710083, + "learning_rate": 2e-05, + "loss": 0.1206, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.545021653175354, + "learning_rate": 2e-05, + "loss": 0.25, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.9636732339859009, + "learning_rate": 2e-05, + "loss": 0.2993, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.263100028038025, + "learning_rate": 2e-05, + "loss": 0.1035, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.4546873569488525, + "learning_rate": 2e-05, + "loss": 0.2404, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.06524017453193665, + "learning_rate": 2e-05, + "loss": 0.2534, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.671353578567505, + "learning_rate": 2e-05, + "loss": 0.4426, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.022323034703731537, + "learning_rate": 2e-05, + "loss": 0.0149, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.4515044689178467, + "learning_rate": 2e-05, + "loss": 0.2001, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.1394161432981491, + "learning_rate": 2e-05, + "loss": 0.1917, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.9682633876800537, + "learning_rate": 2e-05, + "loss": 0.3163, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.6189424991607666, + "learning_rate": 2e-05, + "loss": 0.3118, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.0517044067382812, + "learning_rate": 2e-05, + "loss": 0.1573, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.06075097620487213, + "learning_rate": 2e-05, + "loss": 0.0503, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.1382848471403122, + "learning_rate": 2e-05, + "loss": 0.068, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.777787208557129, + "learning_rate": 2e-05, + "loss": 0.2512, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.3559739291667938, + "learning_rate": 2e-05, + "loss": 0.0825, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.624129056930542, + "learning_rate": 2e-05, + "loss": 0.386, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.7426064014434814, + "learning_rate": 2e-05, + "loss": 0.2817, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.3253116607666016, + "learning_rate": 2e-05, + "loss": 0.2045, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.2575363516807556, + "learning_rate": 2e-05, + "loss": 0.0438, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.630881667137146, + "learning_rate": 2e-05, + "loss": 0.5579, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.7342987656593323, + "learning_rate": 2e-05, + "loss": 0.1118, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.164140224456787, + "learning_rate": 2e-05, + "loss": 0.2386, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.4356032609939575, + "learning_rate": 2e-05, + "loss": 0.1154, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.6298002004623413, + "learning_rate": 2e-05, + "loss": 0.1002, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 4.0662078857421875, + "learning_rate": 2e-05, + "loss": 0.7479, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 5.5125908851623535, + "learning_rate": 2e-05, + "loss": 0.5718, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.290038824081421, + "learning_rate": 2e-05, + "loss": 0.2268, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.18030846118927, + "learning_rate": 2e-05, + "loss": 0.1085, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.576373338699341, + "learning_rate": 2e-05, + "loss": 0.5497, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.286693572998047, + "learning_rate": 2e-05, + "loss": 0.0955, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.9285330772399902, + "learning_rate": 2e-05, + "loss": 0.2454, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.3582700788974762, + "learning_rate": 2e-05, + "loss": 0.3511, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.09600567817688, + "learning_rate": 2e-05, + "loss": 0.2815, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5324358893436928.0, + "train_loss": 0.2733432197570801, + "train_runtime": 129.5029, + "train_samples_per_second": 3.089, + "train_steps_per_second": 0.772 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5324358893436928.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..d267301bf5a839bdca40e76c13bb481172982e73 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f59ece01fe04a7a48f28693d13295bea11713914d282f1437c17f36f13feffd1 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..f1632d409f3780d697826bf5ea9e81d2d1c4270c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23e72f857b9cfaf5796d2873d8235907a34d8882685f3d933ec04def0f03fdb0 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..37c7020e09877a9d85c8dae82718b90e6d03e807 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d4bd14799e8cc7f65cb426d3a0c60dfd4531a12393b8f273347b8c5620ca8bf0 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..93deb17dc66941f73d4dd6c82eadb44645c420cc --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:35f5e34763a844e8c5bedbebb9ce0f4269b6209ce77f42f7dc69997d7d150990 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..70a1e0854f13e962e53030aacc1ca27bbd366225 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d450d5266b0fb3b854b4a8cff0df77291a5d8c573a9f0c8f3e33267f507fa578 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..5874908f6bfdb8563d71bb848a676e0af75d9a0c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dd21af2726d4a658d2b69015fccd86db2148118e4b0fb0ed2c8e0dea5f9e7f8d +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..134a07c8e3c2603c930f58fcb3cfa7ee748a3029 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d31b2c820cabab70ceb4a5441b97331a5bcd878890e2d29c703ce6d5ef4850bb +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..98c01f3966249be785a2909771c5a6a000cbf4fa --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e8bdad959d1e74d007cf25f881664bb0b51b2f0e29db0366edd132e50476b6e +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..cefc16c3d3009a3c6c2ada6bca8a312d0380680d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.656208872795105, + "learning_rate": 2e-05, + "loss": 0.484, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.6125857830047607, + "learning_rate": 2e-05, + "loss": 0.3059, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.17275801301002502, + "learning_rate": 2e-05, + "loss": 0.0445, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.10366298258304596, + "learning_rate": 2e-05, + "loss": 0.0021, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.843791484832764, + "learning_rate": 2e-05, + "loss": 0.6538, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.8289833068847656, + "learning_rate": 2e-05, + "loss": 0.2253, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 4.074177265167236, + "learning_rate": 2e-05, + "loss": 0.331, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.02724658139050007, + "learning_rate": 2e-05, + "loss": 0.0036, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.076066255569458, + "learning_rate": 2e-05, + "loss": 0.1302, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.4889243245124817, + "learning_rate": 2e-05, + "loss": 0.035, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.2523186206817627, + "learning_rate": 2e-05, + "loss": 0.0431, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.04350757226347923, + "learning_rate": 2e-05, + "loss": 0.0025, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 6.011169910430908, + "learning_rate": 2e-05, + "loss": 0.5674, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 8.262720108032227, + "learning_rate": 2e-05, + "loss": 0.3709, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.2983391284942627, + "learning_rate": 2e-05, + "loss": 0.253, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.058121681213379, + "learning_rate": 2e-05, + "loss": 0.0829, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.1454952210187912, + "learning_rate": 2e-05, + "loss": 0.1923, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.8286420702934265, + "learning_rate": 2e-05, + "loss": 0.0425, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.1750544309616089, + "learning_rate": 2e-05, + "loss": 0.0157, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.6171109676361084, + "learning_rate": 2e-05, + "loss": 0.3473, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.4010501801967621, + "learning_rate": 2e-05, + "loss": 0.1544, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.6664732694625854, + "learning_rate": 2e-05, + "loss": 0.5608, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.3054736256599426, + "learning_rate": 2e-05, + "loss": 0.0337, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.1024391651153564, + "learning_rate": 2e-05, + "loss": 0.0633, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.06145457550883293, + "learning_rate": 2e-05, + "loss": 0.1646, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.03543674945831299, + "learning_rate": 2e-05, + "loss": 0.0026, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.007453441619873, + "learning_rate": 2e-05, + "loss": 0.0713, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.725761353969574, + "learning_rate": 2e-05, + "loss": 0.128, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.5994020700454712, + "learning_rate": 2e-05, + "loss": 0.1788, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.03193274885416031, + "learning_rate": 2e-05, + "loss": 0.0039, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.9534406661987305, + "learning_rate": 2e-05, + "loss": 0.2363, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.4525335729122162, + "learning_rate": 2e-05, + "loss": 0.034, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.1786153316497803, + "learning_rate": 2e-05, + "loss": 0.135, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.8209985494613647, + "learning_rate": 2e-05, + "loss": 0.198, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.5794199109077454, + "learning_rate": 2e-05, + "loss": 0.1385, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.22380369901657104, + "learning_rate": 2e-05, + "loss": 0.0329, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.138535022735596, + "learning_rate": 2e-05, + "loss": 0.1659, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.240910530090332, + "learning_rate": 2e-05, + "loss": 0.206, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.9727804660797119, + "learning_rate": 2e-05, + "loss": 0.1167, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.1568636894226074, + "learning_rate": 2e-05, + "loss": 0.0715, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.3930831551551819, + "learning_rate": 2e-05, + "loss": 0.0542, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.7408832907676697, + "learning_rate": 2e-05, + "loss": 0.0493, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.7048364281654358, + "learning_rate": 2e-05, + "loss": 0.0276, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.15886738896369934, + "learning_rate": 2e-05, + "loss": 0.0247, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.0136492252349854, + "learning_rate": 2e-05, + "loss": 0.0828, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.915655255317688, + "learning_rate": 2e-05, + "loss": 0.0251, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.12147674709558487, + "learning_rate": 2e-05, + "loss": 0.0131, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.2046329975128174, + "learning_rate": 2e-05, + "loss": 0.0552, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.31301724910736084, + "learning_rate": 2e-05, + "loss": 0.0181, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.004013196565210819, + "learning_rate": 2e-05, + "loss": 0.0526, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5339656207990784.0, + "train_loss": 0.1446455502510071, + "train_runtime": 138.6509, + "train_samples_per_second": 2.885, + "train_steps_per_second": 0.721 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5339656207990784.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..19232da370e44e842afc8fcbcf129345e77e0776 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c873037bd6c932dc0fb072a7012ca5241802b89f631bd32789cf97b9b3379c14 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..44676d5e752de97f0a2e1a9a295dd3b9d89ffe21 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7689351bd788c6536381fbea755a570ba95ed5ef733a9c2621dc2e02dd0c7e7d +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..6106ac2e329ae281cd3136ee8bee8c541e42ac36 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2bf4cab73754202402344d6e3986f02a70d7f25641b57f59e945d886f00e70dd +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..7e86ef18130e8f8a09658569b3a2e0956e1a6dda --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2c137add24b898892e6223b2640d842b36606f282ad0cf5190408f8a6a7f055d +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..1c1efa682d7d5a4c9ba7b9feb892d964b911bcde --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f86f468da1ab9f11efb57e3d43ec0b8288f742a082cf9e502e3081fc894160ed +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..8cd45bce6609f8485a9eb1560a09090a731ff33e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9582d39a7e1b4e8b7bf7399810318cadd0f24ffa9f7eae78ae5ee3a18f6e8c64 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..cd8e0a848bbbf0091e8e45a61f83a5c70485aef7 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b5f4f500311c3d425019b8534cab769c692435552705b401a849f96f6553df45 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..1656a68f12d30e1fa5157d4494ce9a037215c50b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3b908c756badb69bd67a1b5aea3c62323042db98791b4860de38765456b27d37 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..5a245708e520ab893be2dd427fc1eff96f6790c0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.4923834800720215, + "learning_rate": 2e-05, + "loss": 0.163, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 5.487716197967529, + "learning_rate": 2e-05, + "loss": 0.2793, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 6.219067573547363, + "learning_rate": 2e-05, + "loss": 0.3751, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.7420297861099243, + "learning_rate": 2e-05, + "loss": 0.1123, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.410752773284912, + "learning_rate": 2e-05, + "loss": 0.0934, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.0417581796646118, + "learning_rate": 2e-05, + "loss": 0.1819, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.2629231214523315, + "learning_rate": 2e-05, + "loss": 0.0492, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.5294963717460632, + "learning_rate": 2e-05, + "loss": 0.2905, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.352715015411377, + "learning_rate": 2e-05, + "loss": 0.0843, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 5.852350234985352, + "learning_rate": 2e-05, + "loss": 0.1799, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 8.033079147338867, + "learning_rate": 2e-05, + "loss": 0.3555, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 4.423233509063721, + "learning_rate": 2e-05, + "loss": 0.2215, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.721183776855469, + "learning_rate": 2e-05, + "loss": 0.3109, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.7035958766937256, + "learning_rate": 2e-05, + "loss": 0.2964, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.4034414291381836, + "learning_rate": 2e-05, + "loss": 0.3104, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.549929141998291, + "learning_rate": 2e-05, + "loss": 0.062, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.532846212387085, + "learning_rate": 2e-05, + "loss": 0.3979, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.15882855653762817, + "learning_rate": 2e-05, + "loss": 0.0372, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.4695968627929688, + "learning_rate": 2e-05, + "loss": 0.2874, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.861216068267822, + "learning_rate": 2e-05, + "loss": 0.1257, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.33831530809402466, + "learning_rate": 2e-05, + "loss": 0.0736, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 5.75034236907959, + "learning_rate": 2e-05, + "loss": 0.6081, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.462768793106079, + "learning_rate": 2e-05, + "loss": 0.0732, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.626621961593628, + "learning_rate": 2e-05, + "loss": 0.091, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 7.151983261108398, + "learning_rate": 2e-05, + "loss": 0.6751, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.4828758239746094, + "learning_rate": 2e-05, + "loss": 0.2178, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.1960648149251938, + "learning_rate": 2e-05, + "loss": 0.0197, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.7337711453437805, + "learning_rate": 2e-05, + "loss": 0.0327, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.6478142738342285, + "learning_rate": 2e-05, + "loss": 0.2469, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.7967950701713562, + "learning_rate": 2e-05, + "loss": 0.0823, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 11.077807426452637, + "learning_rate": 2e-05, + "loss": 0.6204, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 6.961493492126465, + "learning_rate": 2e-05, + "loss": 0.1463, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.634577512741089, + "learning_rate": 2e-05, + "loss": 0.2003, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 11.673619270324707, + "learning_rate": 2e-05, + "loss": 0.5573, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.9444699883460999, + "learning_rate": 2e-05, + "loss": 0.4239, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.988490343093872, + "learning_rate": 2e-05, + "loss": 0.1732, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5296922922134399, + "learning_rate": 2e-05, + "loss": 0.1418, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 6.264256477355957, + "learning_rate": 2e-05, + "loss": 0.3325, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.7761623859405518, + "learning_rate": 2e-05, + "loss": 0.049, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.3393197059631348, + "learning_rate": 2e-05, + "loss": 0.1657, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.308199167251587, + "learning_rate": 2e-05, + "loss": 0.1862, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.427755117416382, + "learning_rate": 2e-05, + "loss": 0.1727, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 5.616497993469238, + "learning_rate": 2e-05, + "loss": 0.5919, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 9.243307113647461, + "learning_rate": 2e-05, + "loss": 0.6367, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.188496112823486, + "learning_rate": 2e-05, + "loss": 0.2383, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.5245510935783386, + "learning_rate": 2e-05, + "loss": 0.1459, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.6600544452667236, + "learning_rate": 2e-05, + "loss": 0.0511, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.0679799318313599, + "learning_rate": 2e-05, + "loss": 0.758, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 6.619473934173584, + "learning_rate": 2e-05, + "loss": 0.2778, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.090880870819092, + "learning_rate": 2e-05, + "loss": 0.1221, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2214526502043648.0, + "train_loss": 0.2465058994293213, + "train_runtime": 87.957, + "train_samples_per_second": 4.548, + "train_steps_per_second": 1.137 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2214526502043648.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..b93bee3f327079e11bce68a43859e9646584253a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:644869ddefa49a1209066ec321e4ac5b8968c05dad46998f009cdd825dd60eb6 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..49654b454f65cf65e554f23de79edc9aea8c4097 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63a1cffa46b2a41691e678c39786d3949b9607a7bf0e7cb17530c8a0184b0df9 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..8c80cb735bc8ee46c04b01fc441e578c7f8d44bc --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4aa575295354230f6c2f1aec043601b27afe116beed16e6d0af13cd76c11af0d +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..71fa92c2614fa7c4c3ab7a7f4d0cdafc95bc731f --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a993f2e655003e30118280cf9e19c212d0230c480debc92b1d864ca88ccba8a4 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..75cd04bfd0e181a24e901fc9f13a8275bce3b3f4 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:756aa6c823a65405ce071f2be33adf08abb081f52ca48e93d8021c64c7ef97ca +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b669eaf62ee993043a0a57e2469610ac93a9b8f9 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:faaee8b7a434115a568ab81ebde4929fdd3ab39d1e3d1e5f02bcd9202d10bbc5 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..3d9e69f122853e27bf12acb5edbb569aa510df5c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3db81605f3d3a56406a44cf9f2647f046c43e94ef5ddf0ac5a250b568e2b40b0 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..e15599af1f884c4e92294112a5f433c0d92520fa --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cc71028557d1a4d38657a236dc76e7a30925ebe131c578a14ebb1973d28761f2 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..d9f79a1ca6fce51b790306575834eae542420010 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.726121425628662, + "learning_rate": 2e-05, + "loss": 0.2836, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.31575191020965576, + "learning_rate": 2e-05, + "loss": 0.0352, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 15.00609302520752, + "learning_rate": 2e-05, + "loss": 0.8618, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 7.704536437988281, + "learning_rate": 2e-05, + "loss": 0.232, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.475628852844238, + "learning_rate": 2e-05, + "loss": 0.2153, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.22315561771392822, + "learning_rate": 2e-05, + "loss": 0.0223, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.3340678215026855, + "learning_rate": 2e-05, + "loss": 0.1296, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 6.388670921325684, + "learning_rate": 2e-05, + "loss": 0.3173, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.9941318035125732, + "learning_rate": 2e-05, + "loss": 0.0555, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.18474081158638, + "learning_rate": 2e-05, + "loss": 0.0209, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 8.454588890075684, + "learning_rate": 2e-05, + "loss": 0.6549, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.1090298891067505, + "learning_rate": 2e-05, + "loss": 0.2119, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.6696462631225586, + "learning_rate": 2e-05, + "loss": 0.4728, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.0657601356506348, + "learning_rate": 2e-05, + "loss": 0.1804, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.608464002609253, + "learning_rate": 2e-05, + "loss": 0.5361, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 6.219359397888184, + "learning_rate": 2e-05, + "loss": 0.2275, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.668251097202301, + "learning_rate": 2e-05, + "loss": 0.2561, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 5.5115227699279785, + "learning_rate": 2e-05, + "loss": 0.2005, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.3022406101226807, + "learning_rate": 2e-05, + "loss": 0.1093, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 5.36346960067749, + "learning_rate": 2e-05, + "loss": 0.2803, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.14179377257823944, + "learning_rate": 2e-05, + "loss": 0.0092, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.4682156443595886, + "learning_rate": 2e-05, + "loss": 0.0592, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.8812782764434814, + "learning_rate": 2e-05, + "loss": 0.0907, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 4.291635513305664, + "learning_rate": 2e-05, + "loss": 0.1605, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.8452673554420471, + "learning_rate": 2e-05, + "loss": 0.0455, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.18079137802124, + "learning_rate": 2e-05, + "loss": 0.0755, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.0622549057006836, + "learning_rate": 2e-05, + "loss": 0.0556, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 9.1627836227417, + "learning_rate": 2e-05, + "loss": 0.4496, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.498281955718994, + "learning_rate": 2e-05, + "loss": 0.2297, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.1103051900863647, + "learning_rate": 2e-05, + "loss": 0.0733, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.5733020305633545, + "learning_rate": 2e-05, + "loss": 0.0199, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.2534899711608887, + "learning_rate": 2e-05, + "loss": 0.4748, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.8467509746551514, + "learning_rate": 2e-05, + "loss": 0.0942, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 16.737844467163086, + "learning_rate": 2e-05, + "loss": 1.5072, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.770353078842163, + "learning_rate": 2e-05, + "loss": 0.2233, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 9.26860237121582, + "learning_rate": 2e-05, + "loss": 0.3017, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 8.478885650634766, + "learning_rate": 2e-05, + "loss": 0.8362, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.07585857063531876, + "learning_rate": 2e-05, + "loss": 0.0048, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.762338638305664, + "learning_rate": 2e-05, + "loss": 0.0436, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.6783572435379028, + "learning_rate": 2e-05, + "loss": 0.0386, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 4.228687763214111, + "learning_rate": 2e-05, + "loss": 0.4197, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.5150578022003174, + "learning_rate": 2e-05, + "loss": 0.2686, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 3.67448091506958, + "learning_rate": 2e-05, + "loss": 0.067, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.585147380828857, + "learning_rate": 2e-05, + "loss": 0.2528, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.4488630294799805, + "learning_rate": 2e-05, + "loss": 0.4172, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.668921947479248, + "learning_rate": 2e-05, + "loss": 0.1465, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 8.626337051391602, + "learning_rate": 2e-05, + "loss": 0.7853, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 7.386317253112793, + "learning_rate": 2e-05, + "loss": 0.1106, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.1156115531921387, + "learning_rate": 2e-05, + "loss": 0.1189, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.7104783058166504, + "learning_rate": 2e-05, + "loss": 0.126, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2207748649385984.0, + "train_loss": 0.2561806869506836, + "train_runtime": 84.3718, + "train_samples_per_second": 4.741, + "train_steps_per_second": 1.185 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2207748649385984.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..60fc40af3edf8e11be6b757d25d9d49c8e8bfd12 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:83dbb907fcc7ee236b9e37b9cc29ccf0b752173d7b024722d57a368e60982540 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..98e6fa0772e95797c125967c16025e58951b42dd --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70e452042d4de70a8be7df13b1387c627e1bec0ca08464c090bac60a2d7b3dc9 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..9ec793ec56f081b4d5147302cbc2280820c7f2c0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2baa69058c487b7e33709c8b5a8cc8e03c97f8b241c6e53be7be82e8d33cefcd +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..1270278bc6c6a214c3ccd4344cd368ddd7f2b05a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c668870bc547a5d44b5e76336d6ed818f7bd00c2302c80da0a9eb3d368651b3e +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..60d52c8997ce4110e3615ac12c3f7f3acf50ff74 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a699003d4c0925b3ed677e625f7274ea4d5d6f92db789241944a4bd74769f6cd +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..7dd7bd3e827d3b1f8dfbf9279d5fbce3defea2f0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b254ad8e78f51d19688cbc2343b40d62bcbeccd1d1f8f94ab407b8b5ddbfb8a0 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..973c2336b39949f44a5436cad21e08fb86aae434 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4e2903c7f095a51530beb5f9c5cc799fc68adaa119d93d4704890007ef3f89d9 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d12d6c9f4d9268b76c208180faf7fe57f370c97e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:62de5748d11fb35e5c4c10c0a588885e889120524eed9f49eff5616d45c52e89 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..685f267c1612ab839d8d9c4850dfc9733d6ea197 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.727865695953369, + "learning_rate": 2e-05, + "loss": 0.0586, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.184072971343994, + "learning_rate": 2e-05, + "loss": 0.4171, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 4.215449333190918, + "learning_rate": 2e-05, + "loss": 0.1531, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.0418330430984497, + "learning_rate": 2e-05, + "loss": 0.0545, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 7.7735514640808105, + "learning_rate": 2e-05, + "loss": 0.1424, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.687877655029297, + "learning_rate": 2e-05, + "loss": 0.1672, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.0050475597381592, + "learning_rate": 2e-05, + "loss": 0.0216, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 9.724472999572754, + "learning_rate": 2e-05, + "loss": 0.3274, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 12.362894058227539, + "learning_rate": 2e-05, + "loss": 0.3983, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 11.967941284179688, + "learning_rate": 2e-05, + "loss": 0.888, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.661382675170898, + "learning_rate": 2e-05, + "loss": 0.2159, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 11.347658157348633, + "learning_rate": 2e-05, + "loss": 0.7355, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.9116806983947754, + "learning_rate": 2e-05, + "loss": 0.0817, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.32606440782546997, + "learning_rate": 2e-05, + "loss": 0.0111, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.4042043685913086, + "learning_rate": 2e-05, + "loss": 0.1256, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.5024969577789307, + "learning_rate": 2e-05, + "loss": 0.1121, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.1294269561767578, + "learning_rate": 2e-05, + "loss": 0.0607, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.6410744190216064, + "learning_rate": 2e-05, + "loss": 0.0179, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 5.435573101043701, + "learning_rate": 2e-05, + "loss": 0.3484, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.322148084640503, + "learning_rate": 2e-05, + "loss": 0.2769, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.6120685935020447, + "learning_rate": 2e-05, + "loss": 0.0256, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.1313657760620117, + "learning_rate": 2e-05, + "loss": 0.2728, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.569674491882324, + "learning_rate": 2e-05, + "loss": 0.0616, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 8.987771987915039, + "learning_rate": 2e-05, + "loss": 0.3051, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 3.17429780960083, + "learning_rate": 2e-05, + "loss": 0.3518, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.6236042976379395, + "learning_rate": 2e-05, + "loss": 0.1996, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.2535760402679443, + "learning_rate": 2e-05, + "loss": 0.2249, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 7.700414657592773, + "learning_rate": 2e-05, + "loss": 0.2704, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 9.738489151000977, + "learning_rate": 2e-05, + "loss": 0.3392, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 10.252342224121094, + "learning_rate": 2e-05, + "loss": 0.4548, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 11.407380104064941, + "learning_rate": 2e-05, + "loss": 0.9648, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 7.718284606933594, + "learning_rate": 2e-05, + "loss": 0.7991, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.8864544630050659, + "learning_rate": 2e-05, + "loss": 0.075, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 10.178777694702148, + "learning_rate": 2e-05, + "loss": 0.276, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.7825334072113037, + "learning_rate": 2e-05, + "loss": 0.0562, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.7531514167785645, + "learning_rate": 2e-05, + "loss": 0.4534, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.7974658012390137, + "learning_rate": 2e-05, + "loss": 0.2938, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 4.890337944030762, + "learning_rate": 2e-05, + "loss": 0.7284, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 5.094440937042236, + "learning_rate": 2e-05, + "loss": 0.2009, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.7812694311141968, + "learning_rate": 2e-05, + "loss": 0.1477, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.444733142852783, + "learning_rate": 2e-05, + "loss": 0.2987, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 5.178667068481445, + "learning_rate": 2e-05, + "loss": 0.4161, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 3.907472848892212, + "learning_rate": 2e-05, + "loss": 0.415, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.9277927875518799, + "learning_rate": 2e-05, + "loss": 0.0452, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 5.291507720947266, + "learning_rate": 2e-05, + "loss": 0.5119, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.889734745025635, + "learning_rate": 2e-05, + "loss": 0.1757, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.21428553760051727, + "learning_rate": 2e-05, + "loss": 0.0524, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.8025250434875488, + "learning_rate": 2e-05, + "loss": 0.0601, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.07573585212230682, + "learning_rate": 2e-05, + "loss": 0.0889, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.713853120803833, + "learning_rate": 2e-05, + "loss": 0.059, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2212818824724480.0, + "train_loss": 0.2647566223144531, + "train_runtime": 83.9611, + "train_samples_per_second": 4.764, + "train_steps_per_second": 1.191 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2212818824724480.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..392da8a4a831704ebad3eab3096e1c0adb00a760 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4edc110becb2750243d7ccb3163ccc5c38409132ef51a55b88ee64f5d95511ca +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..9acefceb4d66b28a3dfbc4eabe8be509d0e52515 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8c0039e3e9792220950583df95d2dccbc09487bbcd0b7368825c6fa4d29dfabd +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..f518370fd507cff976faed6535849e903fc9c9bb --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1b67981cc7b51e3bb457b45693efb3a0a6105bbe545e9073b7727b0125fdd7a4 +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..7cb47b6ba1f160744694f02132ff39c2a20c7bf0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e11d9442ab102f5dff212fe05240f80e5d850f8fcff32b68fe37aa60302a419e +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..fec62ed517038aaf67c743574e9cd88e2da7038f --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c84b2bb4436317cd8f476c4f9e802b90f60dab371f74efb5b8978b86dd08c979 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..4c7048345c905826bbf2998d2b83ee48c481e2a1 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bc059e212b196f723caeedd500dbed507682d18c95560bc0aaf32d86d36102ad +size 369839594 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..cb13134247bdc744f93b9dfe1fd566cadc78e77b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:da11b648c5346f06c9a7c3b9a58a95ea2854c69251bf85084a8f1a99735edf51 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..bfc3a5410c1114e0c3cd89a5a6afe35c188ce3df --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:69b6b6d9ee4a85f5c7b0cfc6f846943bc89ac0e5500a07559e354e2ceb1af027 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..a9af38fc773b8b65d56aea070e5806b8d1885110 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.11780647933483124, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.7857205867767334, + "learning_rate": 2e-05, + "loss": 0.0645, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.6256628036499023, + "learning_rate": 2e-05, + "loss": 0.0299, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.09046490490436554, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.1449124664068222, + "learning_rate": 2e-05, + "loss": 0.4328, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 10.677400588989258, + "learning_rate": 2e-05, + "loss": 0.1788, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.7504866123199463, + "learning_rate": 2e-05, + "loss": 0.1158, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.40525200963020325, + "learning_rate": 2e-05, + "loss": 0.0385, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.4561687707901, + "learning_rate": 2e-05, + "loss": 0.0349, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.074843883514404, + "learning_rate": 2e-05, + "loss": 0.1753, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.0829787403345108, + "learning_rate": 2e-05, + "loss": 0.1068, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.774501621723175, + "learning_rate": 2e-05, + "loss": 0.1562, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.8194766044616699, + "learning_rate": 2e-05, + "loss": 0.0823, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 6.974839687347412, + "learning_rate": 2e-05, + "loss": 0.438, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.05832858383655548, + "learning_rate": 2e-05, + "loss": 0.0275, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 9.261703491210938, + "learning_rate": 2e-05, + "loss": 0.2859, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.1834393739700317, + "learning_rate": 2e-05, + "loss": 0.0351, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.24822641909122467, + "learning_rate": 2e-05, + "loss": 0.1897, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.8903193473815918, + "learning_rate": 2e-05, + "loss": 0.0574, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.4879278540611267, + "learning_rate": 2e-05, + "loss": 0.0295, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.6719350814819336, + "learning_rate": 2e-05, + "loss": 0.3643, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.2276092916727066, + "learning_rate": 2e-05, + "loss": 0.0071, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.1784472465515137, + "learning_rate": 2e-05, + "loss": 0.3635, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.444685459136963, + "learning_rate": 2e-05, + "loss": 0.105, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.0479989051818848, + "learning_rate": 2e-05, + "loss": 0.0894, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.3276443481445312, + "learning_rate": 2e-05, + "loss": 0.1098, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 6.08976936340332, + "learning_rate": 2e-05, + "loss": 0.2169, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.46741628646850586, + "learning_rate": 2e-05, + "loss": 0.3431, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.36637449264526367, + "learning_rate": 2e-05, + "loss": 0.0611, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.5141583681106567, + "learning_rate": 2e-05, + "loss": 0.0787, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.01905365101993084, + "learning_rate": 2e-05, + "loss": 0.0178, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.4961856603622437, + "learning_rate": 2e-05, + "loss": 0.131, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.06292741000652313, + "learning_rate": 2e-05, + "loss": 0.1132, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.4960721731185913, + "learning_rate": 2e-05, + "loss": 0.0721, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 7.8650360107421875, + "learning_rate": 2e-05, + "loss": 0.437, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.17741990089416504, + "learning_rate": 2e-05, + "loss": 0.0107, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 5.738003253936768, + "learning_rate": 2e-05, + "loss": 0.1321, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 9.82736873626709, + "learning_rate": 2e-05, + "loss": 0.1997, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 7.798000812530518, + "learning_rate": 2e-05, + "loss": 0.267, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.3873414695262909, + "learning_rate": 2e-05, + "loss": 0.0079, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 5.126296043395996, + "learning_rate": 2e-05, + "loss": 0.2713, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.18566837906837463, + "learning_rate": 2e-05, + "loss": 0.1266, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.7419052124023438, + "learning_rate": 2e-05, + "loss": 0.0119, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.23899133503437042, + "learning_rate": 2e-05, + "loss": 0.0135, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.5020558834075928, + "learning_rate": 2e-05, + "loss": 0.1461, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 8.38952922821045, + "learning_rate": 2e-05, + "loss": 0.1716, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 9.605467796325684, + "learning_rate": 2e-05, + "loss": 0.3758, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.7006304264068604, + "learning_rate": 2e-05, + "loss": 0.0418, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.034456659108400345, + "learning_rate": 2e-05, + "loss": 0.0029, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 7.311607837677002, + "learning_rate": 2e-05, + "loss": 0.2808, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2207554989981696.0, + "train_loss": 0.14105697870254516, + "train_runtime": 89.4687, + "train_samples_per_second": 4.471, + "train_steps_per_second": 1.118 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2207554989981696.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..8c5ded25a927f2b9cea04513b0211684ae2446f0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4f7e09b637ce6d93e754dd5967b190cb4216d0868789655b504a975829bd42e2 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..1a9030890ecbf5eb1e0eb430577a097b93dfb591 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f6d23e0a21ffb7d75ccff9181cdeb9f49f1b184a6c66419ef35438628bbf46ca +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..3c781ad298a0cc4460ea75cf6062551aeeb75396 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7ded3c158f566b8a1dc95032ddc3465bc6eae10e0f64798f6d4861832496a8ad +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b995458295df0445c761b243e690999bbd08bfb1 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:670853ff0783f000df931eb194b28fa580187c204c27cadecd09b52d175fdfa9 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..c9d339a2d376594877c60e96aab2af69df316ac8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:22072483aa2aeb13bff48a71acadb5b606a0e750e3bce69ea642a4cd8717cf51 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..a1dc4c550b99786ce586811a944ef969469f6cee --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fd23173309dac23a4a24dfb4c6747779258b14bda893cb0646659128756c5757 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b25dd2b5542119f7917f9751f8e9718216fd44fa --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7c507e0f7ca80dcf69ed165eb83276ae5d0f4214ee919a734560891498f59e5b +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..36329848905437044863e2011d38fc6beb2d77cc --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:756932425aaa07467c6ae91da9e2d060f359c83aad2f1771d5be5db178b5a351 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3afd4234be070646f03ace10f9288b8b16f785a0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.22604423761367798, + "learning_rate": 2e-05, + "loss": 0.0708, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.35582929849624634, + "learning_rate": 2e-05, + "loss": 0.0647, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.36510753631591797, + "learning_rate": 2e-05, + "loss": 0.0386, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.6686663627624512, + "learning_rate": 2e-05, + "loss": 0.1309, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.5415847301483154, + "learning_rate": 2e-05, + "loss": 0.3175, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.3691963851451874, + "learning_rate": 2e-05, + "loss": 0.0522, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.8830004334449768, + "learning_rate": 2e-05, + "loss": 0.0645, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.5218486785888672, + "learning_rate": 2e-05, + "loss": 0.1094, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.33352574706077576, + "learning_rate": 2e-05, + "loss": 0.3295, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.11427876353263855, + "learning_rate": 2e-05, + "loss": 0.0113, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.0562524795532227, + "learning_rate": 2e-05, + "loss": 0.1144, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.121377944946289, + "learning_rate": 2e-05, + "loss": 0.0514, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.7353308200836182, + "learning_rate": 2e-05, + "loss": 0.1434, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.0296716690063477, + "learning_rate": 2e-05, + "loss": 0.1053, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.3973567485809326, + "learning_rate": 2e-05, + "loss": 0.2621, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.47453567385673523, + "learning_rate": 2e-05, + "loss": 0.4851, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.0953205823898315, + "learning_rate": 2e-05, + "loss": 0.2372, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.4191283881664276, + "learning_rate": 2e-05, + "loss": 0.0905, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2923533022403717, + "learning_rate": 2e-05, + "loss": 0.0406, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.0892467275261879, + "learning_rate": 2e-05, + "loss": 0.0712, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.9888826608657837, + "learning_rate": 2e-05, + "loss": 0.0799, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.672777533531189, + "learning_rate": 2e-05, + "loss": 0.5243, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4991220235824585, + "learning_rate": 2e-05, + "loss": 0.1238, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.28595253825187683, + "learning_rate": 2e-05, + "loss": 0.2601, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.19635741412639618, + "learning_rate": 2e-05, + "loss": 0.2448, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.5289719104766846, + "learning_rate": 2e-05, + "loss": 0.2216, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.3614864647388458, + "learning_rate": 2e-05, + "loss": 0.1776, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.09354989230632782, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.8401954174041748, + "learning_rate": 2e-05, + "loss": 0.1925, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.9142383337020874, + "learning_rate": 2e-05, + "loss": 0.0618, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.11061082035303116, + "learning_rate": 2e-05, + "loss": 0.1339, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.299149990081787, + "learning_rate": 2e-05, + "loss": 0.2535, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.269174098968506, + "learning_rate": 2e-05, + "loss": 0.3316, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.06861578673124313, + "learning_rate": 2e-05, + "loss": 0.0844, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.696437358856201, + "learning_rate": 2e-05, + "loss": 0.3981, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.6215778589248657, + "learning_rate": 2e-05, + "loss": 0.1195, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.2991413474082947, + "learning_rate": 2e-05, + "loss": 0.0269, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.455961227416992, + "learning_rate": 2e-05, + "loss": 0.5541, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.535909414291382, + "learning_rate": 2e-05, + "loss": 0.1846, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.1703977584838867, + "learning_rate": 2e-05, + "loss": 0.1274, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.3546664714813232, + "learning_rate": 2e-05, + "loss": 0.3231, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.05934637039899826, + "learning_rate": 2e-05, + "loss": 0.0218, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.002095094881951809, + "learning_rate": 2e-05, + "loss": 0.0343, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.1383441686630249, + "learning_rate": 2e-05, + "loss": 0.2549, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.6297941207885742, + "learning_rate": 2e-05, + "loss": 0.0307, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.9445375800132751, + "learning_rate": 2e-05, + "loss": 0.0845, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.61324542760849, + "learning_rate": 2e-05, + "loss": 0.0807, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.2835428714752197, + "learning_rate": 2e-05, + "loss": 0.1539, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.11888885498046875, + "learning_rate": 2e-05, + "loss": 0.0359, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.13859808444976807, + "learning_rate": 2e-05, + "loss": 0.2476, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5295755845697536.0, + "train_loss": 0.16337658166885377, + "train_runtime": 128.5136, + "train_samples_per_second": 3.113, + "train_steps_per_second": 0.778 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5295755845697536.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..20a82a1021a9b75f0e5807e5caf4fc28e5ce9663 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:595969ea3df30651670ffe9611706c82d792be2a995d30be6fea6a2aa064bfcc +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..85309ec0a579203c34d454724c633b8385cc12d3 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bfc6648ee2ad9e2dd1223f8faa241563cb462ff07a5581db232d24c48fa37798 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ac8476b34d320e5e3d2d629266fb20322cd5154 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:daa443c70f34f8ef7caa87b03d1e9da75d440510b51aa9fd498a81742247b47e +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..cffa7f7a8fe1989e7bc0ca678c79344568701fa2 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d4768782988f9ec1c4a0b05fa0a6208c81b886784e4348be2a24d12059c5066 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..b2c97c7dec41b24e1e572cb32b33555a9fa98b5a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:26266d91e8fa35964bbceb8ded0c65b1d0048c5431ff687637f84bebc097a470 +size 369837282 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..4448d77fd9d668198eb220ce7c8de976e92e2077 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a0ac977f4fb1b6d5df25bb9765522063f8755ba315582ce0c8a847b170d48eb3 +size 369838470 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..fab6f9808c294eda40b9a88fc8ee449aefe322ff --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8a414876a2690c4b6a0b832245d1e38584cd366c0da78b18fc85106a8272765 +size 369837282 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..70d4922e66a7a7c95205014410b777e802a7e98b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1821499a3d1ce890569577fee932d816d56f0597005e0acfd5ecf57d675f6f57 +size 369837282 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..8da94697b6265fa2150ec535f6f6ac6bebe872d7 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.04902802035212517, + "learning_rate": 2e-05, + "loss": 0.0265, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.07371366769075394, + "learning_rate": 2e-05, + "loss": 0.049, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.2614167034626007, + "learning_rate": 2e-05, + "loss": 0.0047, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.2557220757007599, + "learning_rate": 2e-05, + "loss": 0.0144, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.11752048134803772, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.23466068506240845, + "learning_rate": 2e-05, + "loss": 0.032, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.019244613125920296, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.1971522867679596, + "learning_rate": 2e-05, + "loss": 0.0149, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.1786731481552124, + "learning_rate": 2e-05, + "loss": 0.0616, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.006343411281704903, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.041304152458906174, + "learning_rate": 2e-05, + "loss": 0.0388, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.20497816801071167, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.709629774093628, + "learning_rate": 2e-05, + "loss": 0.1046, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.061534006148576736, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.13374808430671692, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.2791075110435486, + "learning_rate": 2e-05, + "loss": 0.0432, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8878209590911865, + "learning_rate": 2e-05, + "loss": 0.0129, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.0019188154255971313, + "learning_rate": 2e-05, + "loss": 0.0001, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.13442493975162506, + "learning_rate": 2e-05, + "loss": 0.0017, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.8774846792221069, + "learning_rate": 2e-05, + "loss": 0.0127, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.48062625527381897, + "learning_rate": 2e-05, + "loss": 0.0055, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.0540666580200195, + "learning_rate": 2e-05, + "loss": 0.1296, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.01507223304361105, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.03287642076611519, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.011772219091653824, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.04186128452420235, + "learning_rate": 2e-05, + "loss": 0.0403, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.9002057313919067, + "learning_rate": 2e-05, + "loss": 0.0496, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.05484785512089729, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.01304689608514309, + "learning_rate": 2e-05, + "loss": 0.1052, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.01165454089641571, + "learning_rate": 2e-05, + "loss": 0.0082, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.13617275655269623, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.03869406133890152, + "learning_rate": 2e-05, + "loss": 0.0036, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.056844741106033325, + "learning_rate": 2e-05, + "loss": 0.0009, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.11078748852014542, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.009440924040973186, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.0557082891464233, + "learning_rate": 2e-05, + "loss": 0.0105, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.0486905574798584, + "learning_rate": 2e-05, + "loss": 0.0129, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.01112473662942648, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.628478765487671, + "learning_rate": 2e-05, + "loss": 0.0396, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.00634946022182703, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.01175436470657587, + "learning_rate": 2e-05, + "loss": 0.0017, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0016540223732590675, + "learning_rate": 2e-05, + "loss": 0.0002, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.07083069533109665, + "learning_rate": 2e-05, + "loss": 0.0885, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.013444378040730953, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.27856770157814026, + "learning_rate": 2e-05, + "loss": 0.0053, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.1065608263015747, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.015362514182925224, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.451621949672699, + "learning_rate": 2e-05, + "loss": 0.0055, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.05664411559700966, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.04606403037905693, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2217994633609216.0, + "train_loss": 0.01906200110912323, + "train_runtime": 86.4264, + "train_samples_per_second": 4.628, + "train_steps_per_second": 1.157 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2217994633609216.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..7db16e4f3dd7a2b04747c19a8877e6df5528f781 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c1e7afef9b339d67352529d626c9cd4c467e2c3f5673df9d5aa25f6892f830ce +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..c8d050fa2d8b204e28dd265b2a2281ef5ede58b9 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:00ae4c0ab643e8010c4526ac24183e90f32f4c8dd04320560bbe0a3475172151 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..d4e364f8347572d95dd8c529189b4e144d2fc3a5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:28b5a594d4d718cefca7bfc2a4dccf4338bdd5ac1907634e700dfff6433d2518 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..349b2c6e492f23162bb77b5280b0ced1dc0a0a27 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fa08cca9a39a7dc33fdb167682788f4a49678fc1d4b72991d610d35340448443 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..cc6b676521ad9d8fdcfe3066f8c299617dd5f8a1 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fa6d099b54aa680e22dd43048a521aa1e7bd4138a4ef3a56b60749cd0cc43d92 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..2b1bdcf1468c140e48a3d6e46f2ae1b78c8c29df --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2ae24733635b94c4c7bf4b9241c0f084704fe6cc194d52ddd7eaadb3521e27a8 +size 794710050 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..f1202aa7c946b940feb5c90eb3dcde8c73c2005c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:54a75dde28438feed1e447de4e8ed8d15e5e5500cbab63866f4bdeecc995209f +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..aa9f478d8068b7ad01d598664852919b6760d5b6 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5cb748d171d710ae3041a631ce73029ac3247d0d8f6d98bf0b719aef17d93ba6 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..243eaf85648068cb72517b37970d43d7a2d1e810 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.0795845314860344, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.5659167766571045, + "learning_rate": 2e-05, + "loss": 0.0815, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.1368632316589355, + "learning_rate": 2e-05, + "loss": 0.1018, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.605512261390686, + "learning_rate": 2e-05, + "loss": 0.0271, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.03329610824585, + "learning_rate": 2e-05, + "loss": 0.4915, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.11679640412330627, + "learning_rate": 2e-05, + "loss": 0.0073, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.3100295066833496, + "learning_rate": 2e-05, + "loss": 0.0445, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.4292445182800293, + "learning_rate": 2e-05, + "loss": 0.0978, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.671111822128296, + "learning_rate": 2e-05, + "loss": 0.2354, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.21587559580802917, + "learning_rate": 2e-05, + "loss": 0.1476, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.3289199471473694, + "learning_rate": 2e-05, + "loss": 0.0162, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.2584329843521118, + "learning_rate": 2e-05, + "loss": 0.0164, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.0922194719314575, + "learning_rate": 2e-05, + "loss": 0.2975, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.2022584229707718, + "learning_rate": 2e-05, + "loss": 0.0155, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.5872140526771545, + "learning_rate": 2e-05, + "loss": 0.0214, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.5860891342163086, + "learning_rate": 2e-05, + "loss": 0.3171, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.10702989250421524, + "learning_rate": 2e-05, + "loss": 0.0413, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.430460125207901, + "learning_rate": 2e-05, + "loss": 0.0409, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.28363507986068726, + "learning_rate": 2e-05, + "loss": 0.0599, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.5058727264404297, + "learning_rate": 2e-05, + "loss": 0.0686, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.324116051197052, + "learning_rate": 2e-05, + "loss": 0.1618, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.9378544688224792, + "learning_rate": 2e-05, + "loss": 0.0757, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.52661395072937, + "learning_rate": 2e-05, + "loss": 0.1387, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.109260082244873, + "learning_rate": 2e-05, + "loss": 0.0573, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.2523390054702759, + "learning_rate": 2e-05, + "loss": 0.1514, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.4445686936378479, + "learning_rate": 2e-05, + "loss": 0.0183, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.6160315871238708, + "learning_rate": 2e-05, + "loss": 0.1288, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.059535302221775055, + "learning_rate": 2e-05, + "loss": 0.0279, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.017974883317947388, + "learning_rate": 2e-05, + "loss": 0.1993, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.18807987868785858, + "learning_rate": 2e-05, + "loss": 0.0604, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.704691767692566, + "learning_rate": 2e-05, + "loss": 0.2708, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.155722618103027, + "learning_rate": 2e-05, + "loss": 0.5097, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 5.6192307472229, + "learning_rate": 2e-05, + "loss": 0.8215, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.676973581314087, + "learning_rate": 2e-05, + "loss": 0.1603, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.7724549770355225, + "learning_rate": 2e-05, + "loss": 0.0675, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.04755168408155441, + "learning_rate": 2e-05, + "loss": 0.0214, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.12136563658714294, + "learning_rate": 2e-05, + "loss": 0.014, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.48836612701416, + "learning_rate": 2e-05, + "loss": 0.6604, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.441819965839386, + "learning_rate": 2e-05, + "loss": 0.0357, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.015173826366662979, + "learning_rate": 2e-05, + "loss": 0.0257, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.15736602246761322, + "learning_rate": 2e-05, + "loss": 0.0648, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.0816489458084106, + "learning_rate": 2e-05, + "loss": 0.1301, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.21011459827423096, + "learning_rate": 2e-05, + "loss": 0.0267, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.7768332958221436, + "learning_rate": 2e-05, + "loss": 0.2231, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.6184890270233154, + "learning_rate": 2e-05, + "loss": 0.149, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.796445369720459, + "learning_rate": 2e-05, + "loss": 0.3026, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.7821455001831055, + "learning_rate": 2e-05, + "loss": 0.0501, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.05009462684392929, + "learning_rate": 2e-05, + "loss": 0.0032, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.2664127349853516, + "learning_rate": 2e-05, + "loss": 0.1633, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.049664340913295746, + "learning_rate": 2e-05, + "loss": 0.0063, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5296221983866880.0, + "train_loss": 0.1377432131767273, + "train_runtime": 132.5518, + "train_samples_per_second": 3.018, + "train_steps_per_second": 0.754 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5296221983866880.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..2effd92523fa32cde8f8675af6ae6c4aaef8e3a4 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb56910cf864228eab47313f36e15d7c5784a5e034a600a593f48706272ad56d +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..261f7fc28be3c7280a2ce9d387ebdd627713db07 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ec183c2a7eb5bde3af32c648c114654cd5a8e89d7dae1282189fdcf72ea710b2 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..bc5e57aef013507a8eb1b6eab6fe34e4f8e6b97a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dd80b35b4b48150827e1df2c7eb0dadc78a0293b7443aa9d161ebc808c98e1ac +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..20e5197a1eb7cd149d4343fea621d1f2692c694a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a77d5d932b18a15d8606b57ba36462f931f3f9968f1255c97f551478a47cf8c2 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..b0e05a84158224c2416d9969f1003b1493fe4b92 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:411725f2b1e3d592591e6d7b89f32725261e0de0f8967fd967198e724c7fcb29 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6b2754eaf38e97b3ed7208265345c4f84b4e48ba --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1b841bd372c3ed2dde72c503eec05051c94ead5dc0f9b11ea0365ab89ddac173 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..18eb2ddb214cf369491c385586b5e6213193294c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dbc6582ccda3e371503a0aca93bd9b46f7bbfd162191e73bae320677f1de61fb +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..345a3ec994d4a219012715e21f98b1f6bd7be74d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4a25a88714cd1f13980d4fd2320c441c1f05a89da9416f605b4a3db2e3edf9f3 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..6b3c8eb2c2ed6c31b6eb8d7a066268df5a77e147 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.411818265914917, + "learning_rate": 2e-05, + "loss": 0.1544, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.5141104459762573, + "learning_rate": 2e-05, + "loss": 0.6108, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.3105756044387817, + "learning_rate": 2e-05, + "loss": 0.4524, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.03964900970459, + "learning_rate": 2e-05, + "loss": 0.3229, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8813965320587158, + "learning_rate": 2e-05, + "loss": 0.1209, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.9090462923049927, + "learning_rate": 2e-05, + "loss": 0.5076, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.16711433231830597, + "learning_rate": 2e-05, + "loss": 0.175, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.9100112915039062, + "learning_rate": 2e-05, + "loss": 0.2662, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.1325055360794067, + "learning_rate": 2e-05, + "loss": 0.2311, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.6981841921806335, + "learning_rate": 2e-05, + "loss": 0.4229, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.361789345741272, + "learning_rate": 2e-05, + "loss": 0.066, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.655221700668335, + "learning_rate": 2e-05, + "loss": 0.2394, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.399501085281372, + "learning_rate": 2e-05, + "loss": 0.5627, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.3319499492645264, + "learning_rate": 2e-05, + "loss": 0.1707, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.2659320831298828, + "learning_rate": 2e-05, + "loss": 0.0615, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.1162867546081543, + "learning_rate": 2e-05, + "loss": 0.1754, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.5974376201629639, + "learning_rate": 2e-05, + "loss": 0.3857, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.7265806794166565, + "learning_rate": 2e-05, + "loss": 0.0659, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.3086049556732178, + "learning_rate": 2e-05, + "loss": 0.4966, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.758578300476074, + "learning_rate": 2e-05, + "loss": 0.4887, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.9779252409934998, + "learning_rate": 2e-05, + "loss": 0.3501, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.43734505772590637, + "learning_rate": 2e-05, + "loss": 0.0712, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.1445603370666504, + "learning_rate": 2e-05, + "loss": 0.2175, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.46018534898757935, + "learning_rate": 2e-05, + "loss": 0.181, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.0004116296768188, + "learning_rate": 2e-05, + "loss": 0.1226, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.232639789581299, + "learning_rate": 2e-05, + "loss": 0.2379, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.5521408319473267, + "learning_rate": 2e-05, + "loss": 0.2959, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.2255220413208008, + "learning_rate": 2e-05, + "loss": 0.222, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.9854515194892883, + "learning_rate": 2e-05, + "loss": 0.0716, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.10944195091724396, + "learning_rate": 2e-05, + "loss": 0.1438, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.9741544723510742, + "learning_rate": 2e-05, + "loss": 0.3395, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.4257245063781738, + "learning_rate": 2e-05, + "loss": 0.1838, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.6346569061279297, + "learning_rate": 2e-05, + "loss": 0.0739, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.9993082284927368, + "learning_rate": 2e-05, + "loss": 0.2417, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.4346085786819458, + "learning_rate": 2e-05, + "loss": 0.3779, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.5251243114471436, + "learning_rate": 2e-05, + "loss": 0.5186, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.47280192375183105, + "learning_rate": 2e-05, + "loss": 0.042, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.005340814590454, + "learning_rate": 2e-05, + "loss": 0.3701, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.8014842867851257, + "learning_rate": 2e-05, + "loss": 0.0795, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.2434756755828857, + "learning_rate": 2e-05, + "loss": 0.367, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.2541494369506836, + "learning_rate": 2e-05, + "loss": 1.2163, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.14049667119979858, + "learning_rate": 2e-05, + "loss": 0.0182, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.5046660900115967, + "learning_rate": 2e-05, + "loss": 0.3172, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.3243104219436646, + "learning_rate": 2e-05, + "loss": 0.3232, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.07736873626709, + "learning_rate": 2e-05, + "loss": 1.062, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.27489519119262695, + "learning_rate": 2e-05, + "loss": 0.12, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.124854803085327, + "learning_rate": 2e-05, + "loss": 0.2838, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.3364230692386627, + "learning_rate": 2e-05, + "loss": 0.0938, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.619884967803955, + "learning_rate": 2e-05, + "loss": 0.5922, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.3933347165584564, + "learning_rate": 2e-05, + "loss": 0.0647, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5223464331902976.0, + "train_loss": 0.2915180587768555, + "train_runtime": 131.4783, + "train_samples_per_second": 3.042, + "train_steps_per_second": 0.761 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5223464331902976.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..6fe71d22ffdfba240b7e59e65429524de9dee503 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f4636ad6a5923528bb7ec04b09313cad581121bf8a779262ca2b1ef2cfc469c2 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..71e2d12b8ab9da56c15da7c14a82c8a065f8da9b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eedb71a670ce61554d4448d98d7ef7b02955feb74b140a3a22eff6d163768997 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..3abf9afb0db475eb532be798c1aebace5e70ca2c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e351097a97f68887767983497be2f5f7b8b8a58ff139c8120028688504bb4646 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..7fa80c74dede05b9d5e6372bb330adf7a8b1f5a6 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bcecd161a176ca9ec2595768b3206156c3e188157126367ecd2c1b542b572724 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..41402f5f400b8325d265fa35e72e77d529ceb247 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dd72aba5ac2bda58091519b83c4ba5c4f8b6eeef40c4782b00f8bc15e8914975 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..fbe79e98c7242a1441e9b7f4fc1dc04d34fc461b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d01a8a12b2c0c18916d1af4ccd91d3a0e503be7b1ab973b159c3710b04145df8 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c9dbc8443706d398080b658f384567141c29e400 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bdef4f9f711b4fbd743a356b574be88b19569468ca6b5e335f3320556c73d143 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..16fc493377f2edb631111a0a91921acda7c00a0e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:764ab3636ddfd8f9f303420f244893f2f48ea5783bc74b9312df5905ea3c47f1 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..7681c6fdeacd794d76100b7f281a63d0a71f3d64 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.681725263595581, + "learning_rate": 2e-05, + "loss": 0.8213, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.7177492380142212, + "learning_rate": 2e-05, + "loss": 0.3739, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.2289974689483643, + "learning_rate": 2e-05, + "loss": 0.2524, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.5908262729644775, + "learning_rate": 2e-05, + "loss": 0.6541, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.0389397144317627, + "learning_rate": 2e-05, + "loss": 0.5356, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.089683771133423, + "learning_rate": 2e-05, + "loss": 0.7318, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.070936679840088, + "learning_rate": 2e-05, + "loss": 0.6687, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.885744571685791, + "learning_rate": 2e-05, + "loss": 0.5479, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.3766648769378662, + "learning_rate": 2e-05, + "loss": 0.175, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.0584059953689575, + "learning_rate": 2e-05, + "loss": 0.3539, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.8608591556549072, + "learning_rate": 2e-05, + "loss": 0.2411, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.9029541015625, + "learning_rate": 2e-05, + "loss": 0.4995, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.4127941429615021, + "learning_rate": 2e-05, + "loss": 0.232, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.553433656692505, + "learning_rate": 2e-05, + "loss": 0.3638, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.4699469804763794, + "learning_rate": 2e-05, + "loss": 0.4076, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.8907997608184814, + "learning_rate": 2e-05, + "loss": 0.3863, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.6739075183868408, + "learning_rate": 2e-05, + "loss": 0.6892, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.269448757171631, + "learning_rate": 2e-05, + "loss": 0.5923, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.4185779094696045, + "learning_rate": 2e-05, + "loss": 0.206, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.7800968885421753, + "learning_rate": 2e-05, + "loss": 0.4779, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.152641773223877, + "learning_rate": 2e-05, + "loss": 0.6084, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.145423889160156, + "learning_rate": 2e-05, + "loss": 0.6807, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.1414178609848022, + "learning_rate": 2e-05, + "loss": 0.2021, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.6037623286247253, + "learning_rate": 2e-05, + "loss": 0.0641, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.7194721698760986, + "learning_rate": 2e-05, + "loss": 0.5537, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.5481984615325928, + "learning_rate": 2e-05, + "loss": 0.7721, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.06247340515255928, + "learning_rate": 2e-05, + "loss": 0.1662, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.2739735841751099, + "learning_rate": 2e-05, + "loss": 0.2764, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.6438801288604736, + "learning_rate": 2e-05, + "loss": 0.3127, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 3.396937608718872, + "learning_rate": 2e-05, + "loss": 0.5112, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.1666297912597656, + "learning_rate": 2e-05, + "loss": 0.3966, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.8062671422958374, + "learning_rate": 2e-05, + "loss": 0.2849, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.9100825786590576, + "learning_rate": 2e-05, + "loss": 0.5406, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.873441219329834, + "learning_rate": 2e-05, + "loss": 0.8027, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.880889415740967, + "learning_rate": 2e-05, + "loss": 0.3958, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.2626625299453735, + "learning_rate": 2e-05, + "loss": 0.666, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.800750494003296, + "learning_rate": 2e-05, + "loss": 0.6832, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.3463330268859863, + "learning_rate": 2e-05, + "loss": 0.2676, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.7186157703399658, + "learning_rate": 2e-05, + "loss": 0.3298, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.5348477363586426, + "learning_rate": 2e-05, + "loss": 0.167, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.6630174517631531, + "learning_rate": 2e-05, + "loss": 0.128, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.1099162101745605, + "learning_rate": 2e-05, + "loss": 0.2548, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.3519954681396484, + "learning_rate": 2e-05, + "loss": 0.5972, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.1303349733352661, + "learning_rate": 2e-05, + "loss": 0.2681, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5710522532463074, + "learning_rate": 2e-05, + "loss": 0.5044, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.539330005645752, + "learning_rate": 2e-05, + "loss": 0.4028, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.9522987604141235, + "learning_rate": 2e-05, + "loss": 0.2354, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.2203166484832764, + "learning_rate": 2e-05, + "loss": 0.2106, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.2524107694625854, + "learning_rate": 2e-05, + "loss": 0.402, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.6190729141235352, + "learning_rate": 2e-05, + "loss": 0.2314, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5410845953622016.0, + "train_loss": 0.4225331497192383, + "train_runtime": 128.9037, + "train_samples_per_second": 3.103, + "train_steps_per_second": 0.776 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5410845953622016.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..068981d65c7342ae6e397f33e4d7e018de6e0cd0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:172c4ef72648310cd6da903dc28bbfe7bed78acaa6b72ed2fd5ce9b839b74a42 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..b22b852b239ce786263496d930a784ba68acc3cd --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d4cdc0fd045642defb4e1711bfee9b911b6a1be0930dec036ff3d0ef8e36a213 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..b34630cd38bb002f8a872d406eb3ab3530923136 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3f93c94093a4239e3a973fd48188a35a873ff95b8d6a00904fbab27e6b47d94d +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..958629fbaf21fcfba75d9e4e78c7c24aa886a373 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:778a537b9b5c1385c2ddf5333e87c606a50e7b5bd2d4b8b436a1070fab1fc2d5 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..69ffbdce1b9ce3c7f37077f4de8fc234c609337c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:21598cbb558b88944e82ec9d0542c20912cd4d60f0477697a7076e10dcbd73b1 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..914e8becbb1af9b8812aad63d30d9f04e43b04b8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:40e7ab7d92cbd768fbed2b9509d79a18953fa92f0e75b85a57c88ccd04650ec1 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2d3702e2ca5864ce7a331bf1b93548ac3f254d7f --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb6b2f401f1286790d6acd8c869ee1e19ff582358b6ab93053ac2e9aad79c56a +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..cc243c6e87b6deefa65265086036b7a9f6733184 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d6058a5907914723791017de969d0508ab65acfa39a9c4b701b522c2e31d819d +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..144e62ddb58984528e2bcf16aeb4c9c82616ed5a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.36083340644836426, + "learning_rate": 2e-05, + "loss": 0.1761, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.1755833625793457, + "learning_rate": 2e-05, + "loss": 0.3755, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.7400244474411011, + "learning_rate": 2e-05, + "loss": 0.3944, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.1233909130096436, + "learning_rate": 2e-05, + "loss": 0.3499, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8316023349761963, + "learning_rate": 2e-05, + "loss": 0.2026, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.3701503276824951, + "learning_rate": 2e-05, + "loss": 0.1185, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.9007984399795532, + "learning_rate": 2e-05, + "loss": 0.3134, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.1767208576202393, + "learning_rate": 2e-05, + "loss": 0.5549, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.2704261541366577, + "learning_rate": 2e-05, + "loss": 0.168, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.5915101766586304, + "learning_rate": 2e-05, + "loss": 0.1844, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.863276720046997, + "learning_rate": 2e-05, + "loss": 0.9907, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.8704388737678528, + "learning_rate": 2e-05, + "loss": 0.4908, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.4371870756149292, + "learning_rate": 2e-05, + "loss": 0.3113, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.4192326068878174, + "learning_rate": 2e-05, + "loss": 0.4435, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.175323486328125, + "learning_rate": 2e-05, + "loss": 0.3705, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.4722365140914917, + "learning_rate": 2e-05, + "loss": 0.3374, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.6155790090560913, + "learning_rate": 2e-05, + "loss": 0.207, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.142955780029297, + "learning_rate": 2e-05, + "loss": 0.5083, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.7297953963279724, + "learning_rate": 2e-05, + "loss": 0.2052, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.0727624893188477, + "learning_rate": 2e-05, + "loss": 0.6445, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.9438900351524353, + "learning_rate": 2e-05, + "loss": 0.4476, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.977599561214447, + "learning_rate": 2e-05, + "loss": 0.283, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.7979419231414795, + "learning_rate": 2e-05, + "loss": 0.363, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.8474522233009338, + "learning_rate": 2e-05, + "loss": 0.3337, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.1319972276687622, + "learning_rate": 2e-05, + "loss": 0.2235, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.1408110857009888, + "learning_rate": 2e-05, + "loss": 0.6755, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.9507916569709778, + "learning_rate": 2e-05, + "loss": 0.2622, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5303452014923096, + "learning_rate": 2e-05, + "loss": 0.1487, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.3628723621368408, + "learning_rate": 2e-05, + "loss": 0.3303, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.3151333332061768, + "learning_rate": 2e-05, + "loss": 0.4316, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7954598665237427, + "learning_rate": 2e-05, + "loss": 0.2434, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.03599214553833, + "learning_rate": 2e-05, + "loss": 0.4604, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.016643762588501, + "learning_rate": 2e-05, + "loss": 0.3553, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.182877540588379, + "learning_rate": 2e-05, + "loss": 0.3518, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.36678147315979, + "learning_rate": 2e-05, + "loss": 0.4189, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.3178389072418213, + "learning_rate": 2e-05, + "loss": 0.2629, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.47346702218055725, + "learning_rate": 2e-05, + "loss": 0.295, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.4894492626190186, + "learning_rate": 2e-05, + "loss": 0.5294, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.7885500192642212, + "learning_rate": 2e-05, + "loss": 0.373, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.489013671875, + "learning_rate": 2e-05, + "loss": 0.2815, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.7987756729125977, + "learning_rate": 2e-05, + "loss": 0.2361, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.38062584400177, + "learning_rate": 2e-05, + "loss": 0.3566, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.9740220904350281, + "learning_rate": 2e-05, + "loss": 0.1479, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.4188557267189026, + "learning_rate": 2e-05, + "loss": 0.2072, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.9261393547058105, + "learning_rate": 2e-05, + "loss": 0.3558, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.702316403388977, + "learning_rate": 2e-05, + "loss": 0.2526, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.0814781188964844, + "learning_rate": 2e-05, + "loss": 0.3979, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.824862539768219, + "learning_rate": 2e-05, + "loss": 0.3575, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.6064937114715576, + "learning_rate": 2e-05, + "loss": 0.0834, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.408949375152588, + "learning_rate": 2e-05, + "loss": 0.5325, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 6050325026832384.0, + "train_loss": 0.34691070556640624, + "train_runtime": 133.4634, + "train_samples_per_second": 2.997, + "train_steps_per_second": 0.749 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 6050325026832384.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..eee649457edb858f22487f3ceffe0f4a39f49a2a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7b284dd1816538ac40af41b9e922ce85fff296934fdab2e8c9ef5cbccee6fb6b +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..78b5d20427eccc08dc95ee48909e09cc7f90be9c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5493c35a446832190db9dd6d74eb120c49c2f06370639e08f5ccb26474298f20 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..f06992847819d6554e1e85b632a0fb290b42c377 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ad36af6d3d86850165eeb7ff4c03763cf0cfeaa0ddeb5eb132f5d93b63308575 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b24c729d9cdfb3229d0a040fa847a610ff54ff3a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:65d9f1461f1a4414c05b383bddbbe5af38307d233fdf5d3669d2e5dddbfefd39 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..503492bb281f5fc6ce7ddf0f397ed138c3477913 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c623ba12cbc2c64e6fca311a5b669fafccd8b0dba510c3babee4eed4d7ba6b5 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..199202e58a6bb5eb52ef08b1edf99e12182560cb --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aa66a958a49e4f812f465b2dd703a3704645368e7edc0dd9f91d72bce53597ef +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..12a6701d81e2ada7e762198bcf954b18e97e30c5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6304c5ca5b0e0a5f2536ee9000a86acf5c4f1c02ead753b8d1afac2bebabbcd3 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8c9f131fcd4be4395e11df3d37472bfed5f02c03 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8309bf8826b4f62a7cb2c0b930494fa459417fe8714d99b83d4bede82c97e60b +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..a1999957f6f4258a68b4f6bc00f601823e9db4e4 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.5960686206817627, + "learning_rate": 2e-05, + "loss": 0.0804, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.7189565896987915, + "learning_rate": 2e-05, + "loss": 0.2981, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.1532213687896729, + "learning_rate": 2e-05, + "loss": 0.0932, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.2284783124923706, + "learning_rate": 2e-05, + "loss": 0.0825, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.6658846139907837, + "learning_rate": 2e-05, + "loss": 0.1186, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.09774355590343475, + "learning_rate": 2e-05, + "loss": 0.0392, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.005104804877191782, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.09873110800981522, + "learning_rate": 2e-05, + "loss": 0.0237, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.0637214183807373, + "learning_rate": 2e-05, + "loss": 0.1195, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.5293433666229248, + "learning_rate": 2e-05, + "loss": 0.1208, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.7414309978485107, + "learning_rate": 2e-05, + "loss": 0.1093, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.1321890354156494, + "learning_rate": 2e-05, + "loss": 0.1296, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.49531421065330505, + "learning_rate": 2e-05, + "loss": 0.1513, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.4301860332489014, + "learning_rate": 2e-05, + "loss": 0.5732, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.10237313061952591, + "learning_rate": 2e-05, + "loss": 0.0377, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.8299773931503296, + "learning_rate": 2e-05, + "loss": 0.048, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.4566477537155151, + "learning_rate": 2e-05, + "loss": 0.0835, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.19885575771331787, + "learning_rate": 2e-05, + "loss": 0.0344, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.8892322182655334, + "learning_rate": 2e-05, + "loss": 0.036, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.3888661861419678, + "learning_rate": 2e-05, + "loss": 0.1514, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.1144842877984047, + "learning_rate": 2e-05, + "loss": 0.2045, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.9059139490127563, + "learning_rate": 2e-05, + "loss": 0.2374, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.3230063021183014, + "learning_rate": 2e-05, + "loss": 0.1578, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.1620367169380188, + "learning_rate": 2e-05, + "loss": 0.9101, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 5.0908613204956055, + "learning_rate": 2e-05, + "loss": 0.2744, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.226219654083252, + "learning_rate": 2e-05, + "loss": 0.1635, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.07243258506059647, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.789387583732605, + "learning_rate": 2e-05, + "loss": 0.1001, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.15610329806804657, + "learning_rate": 2e-05, + "loss": 0.0118, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.608025848865509, + "learning_rate": 2e-05, + "loss": 0.0351, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.5767613649368286, + "learning_rate": 2e-05, + "loss": 0.0557, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.5605356693267822, + "learning_rate": 2e-05, + "loss": 0.6432, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.4571346938610077, + "learning_rate": 2e-05, + "loss": 0.0567, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.17671281099319458, + "learning_rate": 2e-05, + "loss": 0.0315, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.007368803024292, + "learning_rate": 2e-05, + "loss": 0.2025, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.5134359002113342, + "learning_rate": 2e-05, + "loss": 0.2377, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.07212131470441818, + "learning_rate": 2e-05, + "loss": 0.1799, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.12093732506036758, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.07098981738090515, + "learning_rate": 2e-05, + "loss": 0.53, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.14263422787189484, + "learning_rate": 2e-05, + "loss": 0.0611, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.40487921237945557, + "learning_rate": 2e-05, + "loss": 0.0901, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.10365521162748337, + "learning_rate": 2e-05, + "loss": 0.0563, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.25587397813796997, + "learning_rate": 2e-05, + "loss": 0.0437, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.8094974756240845, + "learning_rate": 2e-05, + "loss": 0.0408, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.7713421583175659, + "learning_rate": 2e-05, + "loss": 0.0263, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.557840585708618, + "learning_rate": 2e-05, + "loss": 0.2098, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.5510768294334412, + "learning_rate": 2e-05, + "loss": 0.0264, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.801114797592163, + "learning_rate": 2e-05, + "loss": 0.2496, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.028664976358413696, + "learning_rate": 2e-05, + "loss": 0.2609, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 6.307191371917725, + "learning_rate": 2e-05, + "loss": 1.3764, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291433615425536.0, + "train_loss": 0.17699591159820557, + "train_runtime": 131.3695, + "train_samples_per_second": 3.045, + "train_steps_per_second": 0.761 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291433615425536.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..9466c070f1563297e818694e83303c91607536d3 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5fef51eb2841095ed227a616f095588592d9bfea157c856175aee1da52b94189 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..fe636cff4b3b0de422a06710067e7a42a6984148 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bbac76399a6aeb264836949cd79b5ab99a06dd25567701beade08af2c1baa1fe +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..04e0bbe5ade0514b7ab70183afc5eb825b9511b6 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d6b8b79914bf00668571022bc6fa5b0a489fac213f703a8de9d29b64784dc36 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..08dc6adb3bd516c824500c4389e85f205946f371 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d58b83f5d4b2276f58f330630d9d1e78761df092a5aaa1284a5c1d714bca8d3 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..dae571c168719bf1d0b9d214cdb5d44c93977a88 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f7d5cc1d89f132e44fd050839c642d8fa66fe150aa0b798cb02396aebd1dbb1 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..cca9c180fa294034990b7e6779a335f46f54a085 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f5379aed128816c70af5d978663d007c96133b230c0ec774bb4d88c899684088 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2e361fb25dcc49dbfaa427e1e191318f5fd43a0b --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fa5fe47587cec705a0ab22ba2092fe20de698bd20cbe495a0bcaedd20262db56 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d694f41bce5b723a506be99789cc25acec532d75 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:473c093fdfb95c2670593088efad2c8f534c97097e7ffbbe153bea46123f0b51 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..bf9bb291ab636eed6a06af6b6b68a96add8ab5da --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.599741220474243, + "learning_rate": 2e-05, + "loss": 0.3678, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.237384557723999, + "learning_rate": 2e-05, + "loss": 0.3552, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.6022919416427612, + "learning_rate": 2e-05, + "loss": 0.3777, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.057524561882019, + "learning_rate": 2e-05, + "loss": 0.4979, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.2288663387298584, + "learning_rate": 2e-05, + "loss": 0.5059, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.5725843906402588, + "learning_rate": 2e-05, + "loss": 0.0434, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.4883065223693848, + "learning_rate": 2e-05, + "loss": 0.2873, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.389308214187622, + "learning_rate": 2e-05, + "loss": 0.5093, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.25661301612854, + "learning_rate": 2e-05, + "loss": 0.7285, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.7146481275558472, + "learning_rate": 2e-05, + "loss": 0.425, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.086560010910034, + "learning_rate": 2e-05, + "loss": 0.3489, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.6659138798713684, + "learning_rate": 2e-05, + "loss": 0.3462, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.4866691827774048, + "learning_rate": 2e-05, + "loss": 0.3146, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.28413724899292, + "learning_rate": 2e-05, + "loss": 0.3458, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.4434587955474854, + "learning_rate": 2e-05, + "loss": 0.9077, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.31303071975708, + "learning_rate": 2e-05, + "loss": 0.3761, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.7123112678527832, + "learning_rate": 2e-05, + "loss": 0.3723, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.1920515298843384, + "learning_rate": 2e-05, + "loss": 0.5283, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.8108996152877808, + "learning_rate": 2e-05, + "loss": 0.3214, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.0596872568130493, + "learning_rate": 2e-05, + "loss": 0.4427, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.1470351219177246, + "learning_rate": 2e-05, + "loss": 0.8642, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7603289484977722, + "learning_rate": 2e-05, + "loss": 0.2048, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.748284339904785, + "learning_rate": 2e-05, + "loss": 0.3988, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.35780930519104, + "learning_rate": 2e-05, + "loss": 0.3286, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.9500956535339355, + "learning_rate": 2e-05, + "loss": 0.4488, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.0172498226165771, + "learning_rate": 2e-05, + "loss": 0.322, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.5446810722351074, + "learning_rate": 2e-05, + "loss": 0.3603, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 3.6236393451690674, + "learning_rate": 2e-05, + "loss": 0.8569, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.191638469696045, + "learning_rate": 2e-05, + "loss": 0.4529, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.2553184032440186, + "learning_rate": 2e-05, + "loss": 0.5221, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.930586576461792, + "learning_rate": 2e-05, + "loss": 0.3677, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.5403631925582886, + "learning_rate": 2e-05, + "loss": 0.4832, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.0856531858444214, + "learning_rate": 2e-05, + "loss": 0.5508, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.274364948272705, + "learning_rate": 2e-05, + "loss": 0.3721, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.8721728324890137, + "learning_rate": 2e-05, + "loss": 0.6982, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.0652120113372803, + "learning_rate": 2e-05, + "loss": 0.3179, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 6.5475687980651855, + "learning_rate": 2e-05, + "loss": 0.8105, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.9676737785339355, + "learning_rate": 2e-05, + "loss": 0.5949, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.5912749767303467, + "learning_rate": 2e-05, + "loss": 0.2847, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.4259419441223145, + "learning_rate": 2e-05, + "loss": 0.2539, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.8830252289772034, + "learning_rate": 2e-05, + "loss": 0.4142, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.9677585363388062, + "learning_rate": 2e-05, + "loss": 0.6066, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.233040690422058, + "learning_rate": 2e-05, + "loss": 0.2817, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.8522083759307861, + "learning_rate": 2e-05, + "loss": 0.5015, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.69613778591156, + "learning_rate": 2e-05, + "loss": 0.3141, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.516747236251831, + "learning_rate": 2e-05, + "loss": 0.9614, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.452475905418396, + "learning_rate": 2e-05, + "loss": 0.2981, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.8674695491790771, + "learning_rate": 2e-05, + "loss": 0.3351, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.743283987045288, + "learning_rate": 2e-05, + "loss": 0.5024, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.030430316925049, + "learning_rate": 2e-05, + "loss": 0.9482, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 1.047142126845952e+16, + "train_loss": 0.4611724853515625, + "train_runtime": 175.0394, + "train_samples_per_second": 2.285, + "train_steps_per_second": 0.571 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 1.047142126845952e+16, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..db52cdf9affbbae50263fd2e13024e7532cd97ce --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:02782480e8aef47226a94c7a60b026d0feb804df3692a145a412bd0ae1e19b07 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..64c69c34df260879fb33622cb1dcf025ee370952 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f6c00fd94152dd2ad8df9cdc38b16f9e33f30e9161b48022860f3cd3e7c2f642 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..2c2bc10292fdd01ee9d70e9a29fb170cbf5ef24d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:625032dec5a2ee81dbcc9fa9d6796c61fa76543ea2e76eba23371501bce32540 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..34da0bfd06985f5c310f0aa20286a91969eb307d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4b21f18652fc468f4ab90cf657b337c789e52ca33661534a87856398a8a464e1 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..22a3fbd65093995808ebcc2e4ba47db525395fcf --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8fadf04d34dcb545f1e70663cf4ef773c1870a556242bd148611e94788b75643 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..c8b4cad9418e2cf5605b9662240dfad2dc16039a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:85df6c9f3b265a035535b492224b97b9a713ac4bf3a423bc8c089455b99f6df0 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..f1e9280dd651b3df7abb19b915e21f7c279d9ba1 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d27f98ab56ac481b628413d62d2f43e623ff577c84bb6af0be8c40d025b24dbd +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..283b5d97193aabccfef3d803a8abf7a2a3c77a1c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7889249f135030fb8194a9368eff6038058bdda590609e7778ad1ab8a9843d3 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..2c57387c210d69a9b85e144daabbac141b898aab --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.9591107368469238, + "learning_rate": 2e-05, + "loss": 0.1289, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.844571828842163, + "learning_rate": 2e-05, + "loss": 0.3757, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.7141473293304443, + "learning_rate": 2e-05, + "loss": 0.1655, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.0768325328826904, + "learning_rate": 2e-05, + "loss": 0.2047, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.050248146057129, + "learning_rate": 2e-05, + "loss": 0.6846, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.12694469094276428, + "learning_rate": 2e-05, + "loss": 0.0603, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.2846057713031769, + "learning_rate": 2e-05, + "loss": 0.0911, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.507358431816101, + "learning_rate": 2e-05, + "loss": 0.384, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.9033132791519165, + "learning_rate": 2e-05, + "loss": 0.188, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.611952066421509, + "learning_rate": 2e-05, + "loss": 0.2868, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.14319366216659546, + "learning_rate": 2e-05, + "loss": 0.2345, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.4576332569122314, + "learning_rate": 2e-05, + "loss": 0.3066, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.5132642388343811, + "learning_rate": 2e-05, + "loss": 0.0296, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.277336597442627, + "learning_rate": 2e-05, + "loss": 0.2683, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.3954222500324249, + "learning_rate": 2e-05, + "loss": 0.2828, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.6852362751960754, + "learning_rate": 2e-05, + "loss": 0.2604, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5023066997528076, + "learning_rate": 2e-05, + "loss": 0.0347, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.4170727729797363, + "learning_rate": 2e-05, + "loss": 0.4175, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.8834854960441589, + "learning_rate": 2e-05, + "loss": 0.0736, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.7064459323883057, + "learning_rate": 2e-05, + "loss": 0.6034, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.5423776507377625, + "learning_rate": 2e-05, + "loss": 0.4638, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.3340253829956055, + "learning_rate": 2e-05, + "loss": 0.2393, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.4756792783737183, + "learning_rate": 2e-05, + "loss": 0.2942, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.19304600358009338, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.4876271188259125, + "learning_rate": 2e-05, + "loss": 0.4545, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.7448945045471191, + "learning_rate": 2e-05, + "loss": 0.4357, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.03668051213026047, + "learning_rate": 2e-05, + "loss": 0.3095, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4883151054382324, + "learning_rate": 2e-05, + "loss": 0.0337, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.1522464007139206, + "learning_rate": 2e-05, + "loss": 0.087, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 6.8636674880981445, + "learning_rate": 2e-05, + "loss": 1.1909, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.3902284801006317, + "learning_rate": 2e-05, + "loss": 0.0507, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.8711813688278198, + "learning_rate": 2e-05, + "loss": 0.5176, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.6282404661178589, + "learning_rate": 2e-05, + "loss": 0.1475, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.6004825234413147, + "learning_rate": 2e-05, + "loss": 0.1389, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.37073016166687, + "learning_rate": 2e-05, + "loss": 0.3907, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.3798085451126099, + "learning_rate": 2e-05, + "loss": 0.2968, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.0137290954589844, + "learning_rate": 2e-05, + "loss": 0.5019, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.38168397545814514, + "learning_rate": 2e-05, + "loss": 0.3715, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.3952118158340454, + "learning_rate": 2e-05, + "loss": 0.3952, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.4785232543945312, + "learning_rate": 2e-05, + "loss": 0.7995, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.565016508102417, + "learning_rate": 2e-05, + "loss": 0.1703, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.3071982860565186, + "learning_rate": 2e-05, + "loss": 0.1136, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.0750744342803955, + "learning_rate": 2e-05, + "loss": 0.4725, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.126516342163086, + "learning_rate": 2e-05, + "loss": 0.3751, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.45165202021598816, + "learning_rate": 2e-05, + "loss": 0.2526, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.7820861339569092, + "learning_rate": 2e-05, + "loss": 0.1652, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.7641480565071106, + "learning_rate": 2e-05, + "loss": 0.0984, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.20326076447963715, + "learning_rate": 2e-05, + "loss": 0.0758, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.0590331554412842, + "learning_rate": 2e-05, + "loss": 0.172, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.9605209827423096, + "learning_rate": 2e-05, + "loss": 0.1944, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5502714666549248.0, + "train_loss": 0.2860134124755859, + "train_runtime": 130.5172, + "train_samples_per_second": 3.065, + "train_steps_per_second": 0.766 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5502714666549248.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..7dd48e98a021d4d99b78bc14b2b4367904cad35a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:69cf82c392ca682a8793122f874802381601e9c591121b65e48598b741d17b1f +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..6035b6facb1e5f528de7994c9d5a1d0a68497f65 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:37dbc1c7a4be0bd0de8d8143199c9d4b4aca843afb0f400c7bb24b1fee0131f3 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..56a9bd208fe70c6347a482c1c755606e2772a193 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b8a6b11acfab620ab9e381816d8c80f8c8f5226dd1fb2c3b4ed80fb15732fc7 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..eea01be86466c824ff88d1407a1bd6aaa9ee6eac --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e4e4da74920bfeb8dacbc823d5fd3ae1f7ce6c9869e2491af822fc94e5a8f4c1 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..d2f08e5990241c7f4d4ea2e16a0ea8bf011b2400 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a22be488d499fcf942508065c5137d084bed22c60a54e79f16c844130777b98f +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..2e782ab5d78fedfb5471c4266397560115238466 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5e157817b64f7ec62a6d9625c544a7e12298778a1e2d18410012d6534ee3b374 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..e5431359f9d21ba75b06c565835e8364c75e5d50 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:14107e14a9faa7d0cf890812d2d75d5ddd26ca17682167139c8515ca166894be +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..eda98d878c1534eb8526e954cac782e982d41605 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:073e58b0fc4729425e1ad5347bd7171c155ac55bbd4fd51ed01a5e427e20f5fa +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..0cf531199d22da5cb87b89e066ad0e1b196af5a6 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.05490674450993538, + "learning_rate": 2e-05, + "loss": 0.029, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.06158506125211716, + "learning_rate": 2e-05, + "loss": 0.0784, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.25062325596809387, + "learning_rate": 2e-05, + "loss": 0.0591, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.901930332183838, + "learning_rate": 2e-05, + "loss": 0.2788, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.04864245653152466, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.175069570541382, + "learning_rate": 2e-05, + "loss": 0.1144, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.047376010566949844, + "learning_rate": 2e-05, + "loss": 0.1069, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 6.073389530181885, + "learning_rate": 2e-05, + "loss": 0.3296, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.5487723350524902, + "learning_rate": 2e-05, + "loss": 0.1173, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.5928170680999756, + "learning_rate": 2e-05, + "loss": 2.1589, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.0531089305877686, + "learning_rate": 2e-05, + "loss": 0.5177, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.5595207214355469, + "learning_rate": 2e-05, + "loss": 0.3314, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.0332772731781006, + "learning_rate": 2e-05, + "loss": 0.127, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.7907194495201111, + "learning_rate": 2e-05, + "loss": 0.0518, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.177078127861023, + "learning_rate": 2e-05, + "loss": 0.0487, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.07242412865161896, + "learning_rate": 2e-05, + "loss": 0.011, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.2309882640838623, + "learning_rate": 2e-05, + "loss": 0.7428, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.759268045425415, + "learning_rate": 2e-05, + "loss": 0.3599, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.079704284667969, + "learning_rate": 2e-05, + "loss": 0.5945, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.1249830722808838, + "learning_rate": 2e-05, + "loss": 0.6052, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 4.680315971374512, + "learning_rate": 2e-05, + "loss": 0.297, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.002690667752176523, + "learning_rate": 2e-05, + "loss": 0.0589, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.320401906967163, + "learning_rate": 2e-05, + "loss": 0.4053, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.5156314373016357, + "learning_rate": 2e-05, + "loss": 0.3833, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.17643266916275024, + "learning_rate": 2e-05, + "loss": 0.2259, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.02747238241136074, + "learning_rate": 2e-05, + "loss": 0.1959, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.29562127590179443, + "learning_rate": 2e-05, + "loss": 0.2263, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.41654130816459656, + "learning_rate": 2e-05, + "loss": 0.2506, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.1793524026870728, + "learning_rate": 2e-05, + "loss": 0.0826, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.08779805898666382, + "learning_rate": 2e-05, + "loss": 0.0117, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.965632915496826, + "learning_rate": 2e-05, + "loss": 0.2234, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.806774854660034, + "learning_rate": 2e-05, + "loss": 0.6122, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.0585837364196777, + "learning_rate": 2e-05, + "loss": 0.0956, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7040287256240845, + "learning_rate": 2e-05, + "loss": 0.0263, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.5498425364494324, + "learning_rate": 2e-05, + "loss": 0.0277, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.4182310104370117, + "learning_rate": 2e-05, + "loss": 0.1829, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.6554776430130005, + "learning_rate": 2e-05, + "loss": 0.0322, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.4884082078933716, + "learning_rate": 2e-05, + "loss": 0.0406, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.14221586287021637, + "learning_rate": 2e-05, + "loss": 0.1089, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.25279563665390015, + "learning_rate": 2e-05, + "loss": 0.0214, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.965247869491577, + "learning_rate": 2e-05, + "loss": 0.4995, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.6208499670028687, + "learning_rate": 2e-05, + "loss": 0.0276, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.1649349331855774, + "learning_rate": 2e-05, + "loss": 0.0088, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.5552606582641602, + "learning_rate": 2e-05, + "loss": 0.0351, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.6708015203475952, + "learning_rate": 2e-05, + "loss": 0.1461, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.175679683685303, + "learning_rate": 2e-05, + "loss": 0.4843, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.22216899693012238, + "learning_rate": 2e-05, + "loss": 0.0076, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.569575071334839, + "learning_rate": 2e-05, + "loss": 0.1574, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.17253278195858002, + "learning_rate": 2e-05, + "loss": 0.0094, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.19309896230697632, + "learning_rate": 2e-05, + "loss": 0.0108, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5329528591220736.0, + "train_loss": 0.2312143099308014, + "train_runtime": 129.9432, + "train_samples_per_second": 3.078, + "train_steps_per_second": 0.77 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5329528591220736.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..2e77b6df18747b93026a769ff842554da118ad24 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:125ac8d851bc7bf80a1a989e79b287c72a2663d91aed4e71dc0dd2c6a40500e8 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..37372435129f3a8193d1a7c468917784bbc2f610 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d3115b14c6b4db856aa25249b819f2b99a6a00e77d725b9296f2feabf01128b +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..33a3ba0c2ef0c881f661de50a4ac3e4cb79eaee6 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ff2f215335c997ae74f94767bfaffabab2c5ae6725f4437e75b249ad441e1486 +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..34117187093d269668a623c269efd26abbb34e3c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:404127fcedca0503be7e4fd2ca50db9eabd9e4b31dec81f40c1ba7f8384423ec +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..7f40ad81530fcc732ee1acd3699cbc6ad9ab6c9e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:59cb85dafc8c8b2421b63abec972460c5aad3455788d553c4e7a4e7783a44fd0 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6bdbdb414a102990ed4fc537ee49036d5f09594d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e25a23fcc7fb4e155a1424fda288949626793e0cd575f669dca8b34441a31d5d +size 794708086 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..e3056322709c03afa5ce0b4061453b77dfd90a8c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f1683175a2a5b7a69d19f645d3f54633e6cf072ee9bfd1a66b870cc409b89468 +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..761d404d2ce0b7f2d1e088aa712f31fba0c20636 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0c2fb08d674259c88ece0369a271b303c9605780c5d6984b402c3875c52cc6ee +size 794706058 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..18c286f0d8a87f8c863fd5c2cd1bd6d1887ee79d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.2572111487388611, + "learning_rate": 2e-05, + "loss": 0.3373, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.44550296664237976, + "learning_rate": 2e-05, + "loss": 0.6416, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.7923219203948975, + "learning_rate": 2e-05, + "loss": 0.0857, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.1195777654647827, + "learning_rate": 2e-05, + "loss": 0.1201, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.14475120604038239, + "learning_rate": 2e-05, + "loss": 0.3056, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.4231327772140503, + "learning_rate": 2e-05, + "loss": 0.1373, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.06305715441703796, + "learning_rate": 2e-05, + "loss": 0.3774, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.551405906677246, + "learning_rate": 2e-05, + "loss": 0.6048, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.5561590194702148, + "learning_rate": 2e-05, + "loss": 0.1561, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.7185628414154053, + "learning_rate": 2e-05, + "loss": 0.2289, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.025214195251465, + "learning_rate": 2e-05, + "loss": 0.7159, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.8984739780426025, + "learning_rate": 2e-05, + "loss": 0.4974, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.3535044193267822, + "learning_rate": 2e-05, + "loss": 0.5082, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.909044623374939, + "learning_rate": 2e-05, + "loss": 0.0666, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.20471078157424927, + "learning_rate": 2e-05, + "loss": 0.18, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.12285207957029343, + "learning_rate": 2e-05, + "loss": 0.0067, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8183131814002991, + "learning_rate": 2e-05, + "loss": 0.5892, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.8510133624076843, + "learning_rate": 2e-05, + "loss": 0.1744, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.28545042872428894, + "learning_rate": 2e-05, + "loss": 0.0584, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.7645012140274048, + "learning_rate": 2e-05, + "loss": 0.3501, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.285703420639038, + "learning_rate": 2e-05, + "loss": 0.2822, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.6761481761932373, + "learning_rate": 2e-05, + "loss": 0.6315, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.5128162503242493, + "learning_rate": 2e-05, + "loss": 0.053, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.2526330947875977, + "learning_rate": 2e-05, + "loss": 0.4234, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.4742984771728516, + "learning_rate": 2e-05, + "loss": 0.0645, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.6137290596961975, + "learning_rate": 2e-05, + "loss": 0.0927, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.6247015595436096, + "learning_rate": 2e-05, + "loss": 0.4699, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.2295064777135849, + "learning_rate": 2e-05, + "loss": 0.0176, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.083214282989502, + "learning_rate": 2e-05, + "loss": 0.362, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.7011772394180298, + "learning_rate": 2e-05, + "loss": 0.0745, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.3976407051086426, + "learning_rate": 2e-05, + "loss": 0.1945, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.3610960245132446, + "learning_rate": 2e-05, + "loss": 0.2118, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.2956864833831787, + "learning_rate": 2e-05, + "loss": 0.1784, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7961469888687134, + "learning_rate": 2e-05, + "loss": 0.0779, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.14717990159988403, + "learning_rate": 2e-05, + "loss": 0.02, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 4.5687947273254395, + "learning_rate": 2e-05, + "loss": 0.4235, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5365333557128906, + "learning_rate": 2e-05, + "loss": 0.1517, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.2739322185516357, + "learning_rate": 2e-05, + "loss": 0.3139, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.5569062232971191, + "learning_rate": 2e-05, + "loss": 0.1229, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.8806869983673096, + "learning_rate": 2e-05, + "loss": 0.3988, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.9513912200927734, + "learning_rate": 2e-05, + "loss": 0.3737, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.31803297996521, + "learning_rate": 2e-05, + "loss": 0.2658, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.40566086769104, + "learning_rate": 2e-05, + "loss": 0.2275, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.7219343185424805, + "learning_rate": 2e-05, + "loss": 0.4791, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.974098801612854, + "learning_rate": 2e-05, + "loss": 0.2705, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.6316423416137695, + "learning_rate": 2e-05, + "loss": 0.2648, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.846888542175293, + "learning_rate": 2e-05, + "loss": 0.5242, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.4156724512577057, + "learning_rate": 2e-05, + "loss": 0.0549, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.9111872315406799, + "learning_rate": 2e-05, + "loss": 0.3027, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.20449677109718323, + "learning_rate": 2e-05, + "loss": 0.0698, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293891360129024.0, + "train_loss": 0.2707864809036255, + "train_runtime": 128.812, + "train_samples_per_second": 3.105, + "train_steps_per_second": 0.776 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293891360129024.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..372076a5034a9fbc922451ae7077f99beac9dffd --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:88f56f6e1de5d8f677d5319c6edf4ee383ef94cda805fa28b2de822afa809ed0 +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..29cce3a2902ef5a6773b19798bf781704eab0b8c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7d99236eaa7dc491f35872d4f39031037edce16b28ab38fa23f4f54efde2755a +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..cf7b088ed84e9edca7e6e6cd9358ce11b58b3151 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cb9b385a09ae9f248f424fef92491de274682cc7382e9a7ac4fd238c54c1b9df +size 571194160 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..038e6ff81f77343f870796580c30e69dcb096ebe --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:11fe32756ef16a4d7ea21707664492646203c3ef711e0b5120058835dd6fa5c6 +size 895441468 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a058278ad4594e1e61753876862206f54f50aa71 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be8ba2b77c174f87711aff5574162137534970e61bb8dbaf002fa6c2d301684b +size 571194160 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..db02dfe019b151e9f37c36d839e5723902379a63 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:18ca098ea991e1a58981e68961b63fb6fb175ed85fdfc18a24ad0f999b8af00c +size 895441468 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9a405f18e1bd4c0148a0b2861d7a09f6a04c59e5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:62cef7e8f38f9b4110aa0e01a6a88df0e101c94c635e4f90965b14a36cc8a501 +size 571194160 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..39d3085b71bcee55cda59ffb45c2268f297a6203 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:32e3134d1aec4a3bc0313d17b375c2d9f135159830169f773394424f473b985b +size 895441468 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0646b55dd604a6ea4617d8b1adbc81bd98bf009e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f056ca728e357ad0ecb7875bfcafb64ad3506d2a403debee80069d0b3699f846 +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e94dc21a0fe488728b358e42a8f3a23df329fedf --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:276eb9632e0321d64e4390547d951a5dcca5dfcdca8b09de376e6fa4fe106e75 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..07d40bc28d3cc527e3bfcb64f97bc10168cca04d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8052465393952005825281ea95edabf74935a2ab99b5e6bd83cc613a0b0c663 +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f7239277bcae9c3acf7792cfa4c7780fe8f4e53f --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:58818cafd851d25502b2f35c5974caaae5311f7a895c4157e4c79dc14c68fecc +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c93fa5ab4e819c7b26c6887d7df7068b2d2491d4 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c072f620ae9581ea11af89c85ce1473a7c9b3009e9599493b596c1f4903565aa +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c545811b751f6121ff2a0aff94738df77af16578 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f852fca7b28b6913b1a38734f6be9f9ec94d700d9ad2e5d9d1bc0a0f175f6ef7 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6f593abfa9e70f5b8b6da2a967bbec33d2d8d258 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:868e0d30db3a0baacdaa697c573d0254887b630f9621a9a25457e7413741e79f +size 571194160 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a04c6b3904d61312cba222fa6a5c968b9985c075 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:abdbe7c01779c642ded44577dc6376e285fae908ab33f61171b6e05d1c336bd4 +size 895441468 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c7f2aec1481a97077d60383f1bc6b6f309854f05 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fbd33e60ef9d4deb2a4682851fe510c094ae4d15c0292e946170fe8d9f34a28a +size 571194160 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d9ded7de8a2edf0e7e3f4c5ec90bd0f053fc5e01 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7a85dbd460a7026becf9cb735403c6a48cdac9db3bccfb49ce8a6bddb9cded69 +size 895441468 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..fa8c8dc20e55024ecad9a7e6c6dc9e1e18c3df58 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:32c85f9a72282eab1abde86805421e41137df9d97939066bf02341cdb9f1ca01 +size 571194160 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b570a35e56118e9f9388d5349974718cb87b7b23 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a72d63021ce3d2c01cb6be06212c767403a1aea65ec1bf1449289d456a6927 +size 895441468 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..83bea2504409394e4e3b3cf3de11b84954574c24 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:01f2c1884a006ea9864bf6200167652c58b23d96694cadde326978c76269cf1d +size 571194160 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..5d2dc6a308528d599604034f38021d34124e87d5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:328f0f9361ccf214c953566da1be893b4a2bcb53ee0ed4c4d6f94e6ac36d9a87 +size 895441468 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..120bf8f00f5492ca141a4c3096f2a685f9f4b303 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d0a277e0beeb6852d1493c417b4412b5bf9fe1d7885143334683ced0cdf31bd +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..812d73327c2d2d1e38efa6e57a37f9e6ad1cdd9c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8c9b86deba9c0210e9921aa99a6031ddc322a07eb1fdac75d66f056adb0bd91f +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..41a3086c81601cd6024f50ebfce6f27fc54dc632 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d44a93b73fabfc71b1e22199c960c806746fbf1eb5711e979713804fcb066542 +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ea51b92b6b2950ce250b965b4cb015175146691f --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:10123c1039b28337399b45cdb4db5622fc706d6b73362b509813ffddf19d74b6 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..bfd7222a71e44731f4ecadadcf5ff94eeef1b879 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d4274de55ab5b32e6c21254a1630ffa15ac252f0c715a5e8e17f8bfab56499bf +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1f4377c0e58f79f3010fe54ea2943d948449fca0 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d1e21d252a711218b3860e60b9adcf69538ac7ed46f806e5d774902bb7d3dcd5 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..bf948dd2c9131b7f41f8f1efcbe2ddb0d63c3500 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29fdf12eb08c896bb51145f17d189975b1c5f6b8dc27ba56f3618e7a09992b2a +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1b1f0b84ee4a10003ac875d1a2c46456aa9e10ef --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:529bd23c95b9e0d1f8f9a226e11ec85fb5a60ba17628cf59c7ef3a82164a9741 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4d1d149fc432af751a5ab8957815f26773ac06e5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23f17ced089b79b59a904223e0fa1fe12fa54c6c29d2c086f58097a5b5f31fcd +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..31408381ecce6a71693085bd81a6bb36d66fdf01 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:28cd389f18c5675322aa66e5f631ede2f8ce9450badb5fae982070fd05d46ab6 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..873addc848d577bce34da49c00120d3a7ab05f99 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9bf252418a70c9f6328c8add87b5d36a354ff927053d9c75eb26e4f36cd1870e +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1254482f8e727482e6dce7587be30c265b5b4e5c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a020100d0e20461ea7be48fc4d8cf1e1d286f3bcef87680cbae1e4be5a5ffc20 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f473e21007190a0571e804eaea8719923b66faae --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1ceb166969e858dd9f14176b3aed096864575f21943925f22f366f91a120321c +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6275f131fd778084b6bdc3a01b62b80218031635 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1248fc034d643a0e5e2e32024703ee4f022cf8a6acb202242c05345cafc628ce +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4fdf724df19a3f3f4a51f74c67566c745d15b88c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:91b26ed37563e643eeb0cdb1bccf4cc05d7ddd9515275d39f52a86b6548fc5ba +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..fd918d248554d948e9bbf16ccccec77c29ddc3a3 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2a00210ed16cd6f7274eac0d313435bfac97cf975294688824f10fe61fa27eb2 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b2dc75a1c031502f40aea0f1db36caab843bbb7d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1ce028873029aea94a8b7b3a141ba33e220cc33452aba0cc7b370ea2639e318e +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a7fc608f384eddbf348a01a8a5b3601c98d21d95 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2e8a28107a7ea3fdc9d8715bd1d46143a3d775cb9e76d6515ebe5e915ce783f0 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f66356ec5ded0762af05b88b6f409562f20fd72e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:73b7258d14eb63fb12801a33b8f9ad10e8cb09c03386e84fc0ecbb18b58b00ae +size 1724097712 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a21c81fc57a82e8e400d372fb8ca97a5e4de0b5f --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:281ed28285a9423052bbaa74b977b71c0433a53af497163a6081e0a64647d5b5 +size 1076152836 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/latest b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..8480f3290fe02dcc8ad3f411dce4bb771607f136 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2a7a8d47edc80b205a0230dd0d312c16f3cc81aa0b01bd0d21718b6314735bb +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..120369764554a4d6cac58078d8512bf359607829 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf1dde35905b1ad125edf778aff1f9c3576ae0715c8c9876f665a63efcf18cd +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..64b454d945aa2a842c509c5cebe5ce0706d83a0c --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb23ed1efc2080025ddd7ca376ff771b3f058afa24e048f5cb3cf99201aa1e03 +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ad8d07954c41361e73a0c9c7d3b4b46d894b188 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8caca9dd439d57dbc5054f65d33ebbd7ac694940b2b9dbe5e8cdf68881137a2 +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f05da1a0577050cea961f470b13842ab5f1d3f9 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafca5706ac84f575d7648ee4786543561523e58cea8c01d2214b03d300bb988 +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..351a4d2a284717aec90435f50067e1bb7002e9e4 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:341a5a225c4c17a621d4f60caf4c34834000ea74e6f3c8b4f5aaada72ca1f547 +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..08cfb19e58b68e297c605b87165d4acd241054e5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb378347913e2dd7be50fa6f9e3bd9e481f3897c2c9d31c597fe9164efc50bee +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..aeee08397c6123d33648f6703464379a4b7f85fe --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70c677ae195453774b8815440fee856ddb773c322681ac422ebc6e8abd5a58d7 +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9935621fdc6da9e4abad379d97303b5ebf4f8f5 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:393c3d5a15b40b2ff049c254f5b09f25d226b0971827b4e7460ca58f9ff7b76f +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..2766cf1a4da1390df01ea4146b4a3a7a06e1434e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dc77db94ac5b4746eebf2c850a9425da9a425a71108790d6b7f5f2654b614a8 +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..34bfc5121e29fc587c3f574c17d8006189d3021a --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47a94b88b812f21101345aae20a53cc10e0b346660c2cef669c77b201d06d508 +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..4be93ba48f85d89c8cbd0e7426b79a10f51bc786 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20b9a2e67fdcb4051cebf52cba2ed0d47facb4dc69170f35a3483e3601bf25a4 +size 704650017 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..b42d717305b30cd186c1c408b8719020c0658a5d --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fd7458877767da0a36b80344cd6ec7eae8fbc8e11d07f8bfce3b4d5f5f239a9 +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..01842d1316f21c3326c5735c6257b499da7ca458 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e537f056156082569ae0819abd06beb04b46581d0c63843bd68926e0ac2bff0e +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..3af9fe115499486166c3ef6d0a7bd27661ac08b4 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a9649d6997f84508e5f566393e91e42e03a31aca55878bacd7ffd2e134bbfc +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..abd1e557c23a8769870b0beedaf7df93836a101e --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7520bace4f501771d879621facd1b44f8db1ad7385a653e1535c7642465e67d +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..ebc846b7a5c587ced387cca85c4065e921c8cbe3 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1951a910d1b7587459db2110e8ac43cbc906a34267372eb21d2244a900eaec46 +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..5eb1370d3fa3177cb4fb710882f4b7eb706817e8 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03f40a0df2437812b9e6de4e5d4e30014acb71344efa8132f19e1fc0f5925c82 +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..385c599585a0973f240a027d28e985ca72d6b721 --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3ef6d8247bc37b017895c9e5bbaf32f2db500f94f66b4ebade5a4a5ba7a271 +size 704649992 diff --git a/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..219adb8927937a4597cd82efbf7ebb4c8a7519ad --- /dev/null +++ b/client_states_feddat_feddualMulti2pqfullfreeze_back_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:277fa5f0607d993e80153171e4760f999c4ee6a65452ea62694d0b9b03818f22 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..c14580953fc97afe9fd1707dc6855bb7b40c1366 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf648ad1b7f919da8a2a68acbd6b733b20af21e5e39576f60fc897201e1f96de +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..b57cbe698bf5de29b45c931916370b12cb1feb8e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3be98eae9a38c00393429970583417db7505a8e6c780f7aae49bbf0aa247e0ff +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..dbd5ab6a3d3e9c94de6e64ef7cc5dd33803bfc20 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7067e8fc2125589a6e3ca521790acad3526025b4809d0999acbb80664d72022a +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..f4815a374e0f98325eb2debe9b68ac159c86b4f0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b88b0f417cb30f390db8b01f44d152338e9003a933fc17f3983de16f712c9013 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5316187f5ac16781bd6fa3dc4def1856baa10f13 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:04e05efbabbc09b1b4e385dfe696e686f510897398857bd4e30122e52735e8ab +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..bd60cffcde12d17232f8936ed85f6ba2e94b3645 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8c10bbc53cf4e32f0e9635d1a56ee1e92171aebc1e31b9dcc1e8a4673d20f24b +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..eca2440511581370d40b6c29ae9ee4c5c255bdca --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a476a84011990407a46d77a097e1f143d986220ce71f9d3e110a36c753af5cdf +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..aa8e73438cd76a0a05b14756823bd152f4ad4b76 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:316ffadb2823e21b6e3b1877b6c2a5d0f19ecf26550c05a7f8ce6285c6b95d2c +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..7127a7bc8d26c605e7dd8826a8ec3a53b04ff009 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/0_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.753019094467163, + "learning_rate": 2e-05, + "loss": 0.1975, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.3112585544586182, + "learning_rate": 2e-05, + "loss": 0.2376, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.5298196077346802, + "learning_rate": 2e-05, + "loss": 0.0317, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.22096067667007446, + "learning_rate": 2e-05, + "loss": 0.0355, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.5044991970062256, + "learning_rate": 2e-05, + "loss": 0.2437, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.51603102684021, + "learning_rate": 2e-05, + "loss": 0.251, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.0583287477493286, + "learning_rate": 2e-05, + "loss": 0.0855, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.1202094554901123, + "learning_rate": 2e-05, + "loss": 0.5048, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 5.4441022872924805, + "learning_rate": 2e-05, + "loss": 1.0054, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.0599565505981445, + "learning_rate": 2e-05, + "loss": 0.2779, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.13993246853351593, + "learning_rate": 2e-05, + "loss": 0.0129, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.7802276611328125, + "learning_rate": 2e-05, + "loss": 0.3702, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.18995238840579987, + "learning_rate": 2e-05, + "loss": 0.2015, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.2972745895385742, + "learning_rate": 2e-05, + "loss": 0.3064, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.683105230331421, + "learning_rate": 2e-05, + "loss": 0.297, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.11963244527578354, + "learning_rate": 2e-05, + "loss": 0.0263, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.894568681716919, + "learning_rate": 2e-05, + "loss": 0.153, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 5.43541955947876, + "learning_rate": 2e-05, + "loss": 0.3892, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.7024756073951721, + "learning_rate": 2e-05, + "loss": 0.2805, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.6221795082092285, + "learning_rate": 2e-05, + "loss": 0.1208, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.2980841398239136, + "learning_rate": 2e-05, + "loss": 0.0228, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.6763508319854736, + "learning_rate": 2e-05, + "loss": 0.3016, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.7829555869102478, + "learning_rate": 2e-05, + "loss": 0.1275, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.25025421380996704, + "learning_rate": 2e-05, + "loss": 0.2677, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.430356740951538, + "learning_rate": 2e-05, + "loss": 0.0688, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.30621075630188, + "learning_rate": 2e-05, + "loss": 0.0988, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.4240196645259857, + "learning_rate": 2e-05, + "loss": 0.1184, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.154081106185913, + "learning_rate": 2e-05, + "loss": 0.134, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 4.324033260345459, + "learning_rate": 2e-05, + "loss": 0.6193, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.11481756716966629, + "learning_rate": 2e-05, + "loss": 0.0068, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.6372262239456177, + "learning_rate": 2e-05, + "loss": 0.1449, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.6423384547233582, + "learning_rate": 2e-05, + "loss": 0.0927, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.4643678069114685, + "learning_rate": 2e-05, + "loss": 0.0183, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.1954665184020996, + "learning_rate": 2e-05, + "loss": 0.0922, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.5968165397644043, + "learning_rate": 2e-05, + "loss": 0.2534, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.9025938510894775, + "learning_rate": 2e-05, + "loss": 0.0975, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 5.357425689697266, + "learning_rate": 2e-05, + "loss": 0.8433, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.2006731033325195, + "learning_rate": 2e-05, + "loss": 0.3677, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.678015947341919, + "learning_rate": 2e-05, + "loss": 0.282, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.619877278804779, + "learning_rate": 2e-05, + "loss": 0.2936, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.4968724548816681, + "learning_rate": 2e-05, + "loss": 0.0823, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.6745550632476807, + "learning_rate": 2e-05, + "loss": 0.2108, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.5290979743003845, + "learning_rate": 2e-05, + "loss": 0.0302, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.4467601776123047, + "learning_rate": 2e-05, + "loss": 0.1418, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.3250622749328613, + "learning_rate": 2e-05, + "loss": 0.1212, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.5787934064865112, + "learning_rate": 2e-05, + "loss": 0.2072, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.844336986541748, + "learning_rate": 2e-05, + "loss": 1.0388, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.6109411716461182, + "learning_rate": 2e-05, + "loss": 0.4332, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.78183650970459, + "learning_rate": 2e-05, + "loss": 0.3449, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.29656916856765747, + "learning_rate": 2e-05, + "loss": 0.0665, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289495989583872.0, + "train_loss": 0.23913730621337892, + "train_runtime": 123.4957, + "train_samples_per_second": 3.239, + "train_steps_per_second": 0.81 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289495989583872.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..d1c2435b17b01770c6ae3ed3a04bb18eab9b476c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c097133822ce54995f79abf7a9aac5e2f95cd66ed0dbeb7cba740fc59a0a4b44 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..509dc665c0669c468131b338e4492fe71c3df406 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4480d521e5e1ef36cd7bb7c031e28f4c2ba4ae0aa6173daece37b23093b30d0b +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..978aeceda25071147189738f54038caef8a059c3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aa1a23c19b9a2d022800273df27a04f9c1d2752dcf9240400b606f4b4ea012f3 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..9492715a9fff70edd600cae179b7e3aa6745bac1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:755aebcd63d08144d8ba34377559427c8bcdf7fd5fd3456e5c46b095b8c37460 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..aa48255b579dc27f90af71ff1d61c13534c5407c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4eeed1cd892d7d2f5385512bca1fd510f6872bb02eec1b4b1f40ecdba0946709 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..2fc59b62e5062369d0e338a78e9fc01cfb5ae50e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1a215f5c685cafadf62611a20d9eee9218dad791377c70f0e314b74f5dee41ee +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2dd402a500dc167291805b7ccb8f1df19ac19bba --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:12bf1cd889bf1542952978d292b3e4f7b05a60e959a8f55036ce0a5e1c75eba3 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8ce01773b20533ab1e52f7e1f2ad79ea978845c4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:074f863bd05caaefb2b7999931af1c79d04277c6f95adb5953f5fd848e233286 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..e9b96464bbb7f52498f90ff6d0aa1b8dd6e8e70b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/10_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.558692932128906, + "learning_rate": 2e-05, + "loss": 0.2455, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.388980865478516, + "learning_rate": 2e-05, + "loss": 0.2039, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 4.012956142425537, + "learning_rate": 2e-05, + "loss": 0.1342, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 5.421528339385986, + "learning_rate": 2e-05, + "loss": 0.6346, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.7387746572494507, + "learning_rate": 2e-05, + "loss": 0.1909, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.1450187861919403, + "learning_rate": 2e-05, + "loss": 0.0121, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 8.148824691772461, + "learning_rate": 2e-05, + "loss": 0.493, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.4360415935516357, + "learning_rate": 2e-05, + "loss": 0.2185, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 5.9994964599609375, + "learning_rate": 2e-05, + "loss": 0.1761, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.4697868824005127, + "learning_rate": 2e-05, + "loss": 0.2366, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 5.762057304382324, + "learning_rate": 2e-05, + "loss": 0.2459, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.399719476699829, + "learning_rate": 2e-05, + "loss": 0.2024, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 6.002265930175781, + "learning_rate": 2e-05, + "loss": 0.5301, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 12.767424583435059, + "learning_rate": 2e-05, + "loss": 0.6303, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 11.84737777709961, + "learning_rate": 2e-05, + "loss": 0.5659, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.0040841102600098, + "learning_rate": 2e-05, + "loss": 0.0467, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.0845627784729004, + "learning_rate": 2e-05, + "loss": 0.1099, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.16188879311084747, + "learning_rate": 2e-05, + "loss": 0.2984, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 6.404210090637207, + "learning_rate": 2e-05, + "loss": 0.316, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.038934707641602, + "learning_rate": 2e-05, + "loss": 0.2457, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.7313228249549866, + "learning_rate": 2e-05, + "loss": 0.1023, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.14477340877056122, + "learning_rate": 2e-05, + "loss": 0.639, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.445493459701538, + "learning_rate": 2e-05, + "loss": 0.0639, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.8347914814949036, + "learning_rate": 2e-05, + "loss": 0.2348, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.174941301345825, + "learning_rate": 2e-05, + "loss": 0.0955, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.3765628337860107, + "learning_rate": 2e-05, + "loss": 0.6842, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 13.489584922790527, + "learning_rate": 2e-05, + "loss": 1.1299, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.3781820237636566, + "learning_rate": 2e-05, + "loss": 0.0146, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 6.7035908699035645, + "learning_rate": 2e-05, + "loss": 0.1161, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 10.01025390625, + "learning_rate": 2e-05, + "loss": 0.3398, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 5.989887714385986, + "learning_rate": 2e-05, + "loss": 0.364, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.107100486755371, + "learning_rate": 2e-05, + "loss": 0.0844, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.42985200881958, + "learning_rate": 2e-05, + "loss": 0.9404, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 4.179809093475342, + "learning_rate": 2e-05, + "loss": 0.3732, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.9546871781349182, + "learning_rate": 2e-05, + "loss": 0.0292, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 19.416044235229492, + "learning_rate": 2e-05, + "loss": 1.041, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.932291507720947, + "learning_rate": 2e-05, + "loss": 0.221, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 8.879201889038086, + "learning_rate": 2e-05, + "loss": 0.2715, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 6.588296890258789, + "learning_rate": 2e-05, + "loss": 0.8101, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 5.008404731750488, + "learning_rate": 2e-05, + "loss": 0.8212, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 8.52298355102539, + "learning_rate": 2e-05, + "loss": 0.4807, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.09916418045759201, + "learning_rate": 2e-05, + "loss": 0.0057, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 5.215274333953857, + "learning_rate": 2e-05, + "loss": 0.322, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 6.467981338500977, + "learning_rate": 2e-05, + "loss": 0.6716, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.1091408729553223, + "learning_rate": 2e-05, + "loss": 0.1458, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.7563092708587646, + "learning_rate": 2e-05, + "loss": 0.267, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.696639537811279, + "learning_rate": 2e-05, + "loss": 0.7468, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.233278751373291, + "learning_rate": 2e-05, + "loss": 0.1309, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.206926345825195, + "learning_rate": 2e-05, + "loss": 0.2918, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.7804235219955444, + "learning_rate": 2e-05, + "loss": 0.563, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2220576018006016.0, + "train_loss": 0.35476917266845703, + "train_runtime": 66.5811, + "train_samples_per_second": 6.008, + "train_steps_per_second": 1.502 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2220576018006016.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..dd31a8f382e06aa1c3ebaa8dfdee6cdd266985de --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6cc7922aed2ca879e2e8856dc153a7018cec0d3671a074993e17b80472ff8794 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..34808710067db8b0bee7dd0b61ff96e49e3c06d7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e43ed5bbbe370ec6c554180543acef7e498ea7d8682cc2d89ec2dd4aea84156b +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..f99c1e13cfc86dc7e2084eda64c9bb3ab69016cc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ecd7d20e45a5d34befd4e2ee64bc713d6ef98c6b3b226370fa39059a99f26fa1 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..cbd3f0419b33edaa46ea76e9b0deaa83801ddace --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb73f286273c275450d4406e0ffe8163c9c727d85d547100c7195fdbf6842f94 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..790be7ab21e393dc350e9f95d17147991c71851b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b398adac3eb2ff92b020e70acfbde76d9f601e7ce289217a288b29b47711c1b9 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..0a1d4444b0fc5a3906efe4aa6dcf6eaa37eff39b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9d1a1b99d25eaa0aacebe043974cb4bc7c2065e0a3a5318b212cb11ad1f53c20 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..87c8db3cd218071085bca80549a60666ffe358b0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ee8e3fa71eafc1369ac4acfa7e29ffe0155a4d139445f946d082400a7fad581d +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..e1d53acaff192846b42e9cf2b67d684d2f55372b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:939e7c40e3a1c5d2997471406e5df01e89850c594ab8c49efb6c1feaaf78cea5 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c4a318861f21a41ff165b557b8723069354e56c5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/11_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 5.078338146209717, + "learning_rate": 2e-05, + "loss": 0.5278, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.5660552978515625, + "learning_rate": 2e-05, + "loss": 0.3679, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.4088938236236572, + "learning_rate": 2e-05, + "loss": 0.4324, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.3407633304595947, + "learning_rate": 2e-05, + "loss": 0.2724, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.688930511474609, + "learning_rate": 2e-05, + "loss": 0.8701, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.924493432044983, + "learning_rate": 2e-05, + "loss": 0.4498, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 4.788064479827881, + "learning_rate": 2e-05, + "loss": 0.9287, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 6.805437088012695, + "learning_rate": 2e-05, + "loss": 0.7769, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.5954489707946777, + "learning_rate": 2e-05, + "loss": 0.6084, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.7468198537826538, + "learning_rate": 2e-05, + "loss": 0.2982, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.0640840530395508, + "learning_rate": 2e-05, + "loss": 0.3046, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.8069891929626465, + "learning_rate": 2e-05, + "loss": 0.4462, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.5049664974212646, + "learning_rate": 2e-05, + "loss": 0.3929, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 7.184503555297852, + "learning_rate": 2e-05, + "loss": 0.459, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.6318819522857666, + "learning_rate": 2e-05, + "loss": 0.5269, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.5338833332061768, + "learning_rate": 2e-05, + "loss": 0.4568, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 4.059306621551514, + "learning_rate": 2e-05, + "loss": 0.5347, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 5.128389835357666, + "learning_rate": 2e-05, + "loss": 0.4922, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.5832765102386475, + "learning_rate": 2e-05, + "loss": 0.5139, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.7789223790168762, + "learning_rate": 2e-05, + "loss": 0.4738, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 6.092098236083984, + "learning_rate": 2e-05, + "loss": 0.3383, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 6.010416030883789, + "learning_rate": 2e-05, + "loss": 0.6818, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.193341016769409, + "learning_rate": 2e-05, + "loss": 0.5159, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.882288932800293, + "learning_rate": 2e-05, + "loss": 0.2616, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 5.338334560394287, + "learning_rate": 2e-05, + "loss": 0.6005, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.119640827178955, + "learning_rate": 2e-05, + "loss": 0.5308, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.219879150390625, + "learning_rate": 2e-05, + "loss": 0.3052, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 3.78957200050354, + "learning_rate": 2e-05, + "loss": 0.2678, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.9426381587982178, + "learning_rate": 2e-05, + "loss": 0.25, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.653050184249878, + "learning_rate": 2e-05, + "loss": 0.3029, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.747711658477783, + "learning_rate": 2e-05, + "loss": 0.3892, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.814879894256592, + "learning_rate": 2e-05, + "loss": 0.3433, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 5.769069194793701, + "learning_rate": 2e-05, + "loss": 0.604, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.09402260929346085, + "learning_rate": 2e-05, + "loss": 0.1638, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.9431726932525635, + "learning_rate": 2e-05, + "loss": 0.3768, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.6337473392486572, + "learning_rate": 2e-05, + "loss": 0.3818, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.390547513961792, + "learning_rate": 2e-05, + "loss": 0.5308, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.0507762432098389, + "learning_rate": 2e-05, + "loss": 0.4029, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 5.858564853668213, + "learning_rate": 2e-05, + "loss": 0.6493, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.9529476165771484, + "learning_rate": 2e-05, + "loss": 0.5532, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.248067855834961, + "learning_rate": 2e-05, + "loss": 0.7119, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 4.329884052276611, + "learning_rate": 2e-05, + "loss": 0.708, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.44238758087158203, + "learning_rate": 2e-05, + "loss": 0.3341, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.5500311851501465, + "learning_rate": 2e-05, + "loss": 0.4734, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.027665594592690468, + "learning_rate": 2e-05, + "loss": 0.4622, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.270262241363525, + "learning_rate": 2e-05, + "loss": 0.3262, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.0984396934509277, + "learning_rate": 2e-05, + "loss": 0.4897, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.927386522293091, + "learning_rate": 2e-05, + "loss": 0.4346, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.2815515995025635, + "learning_rate": 2e-05, + "loss": 0.4052, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 6.492791175842285, + "learning_rate": 2e-05, + "loss": 0.4614, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2191347477905408.0, + "train_loss": 0.4678009796142578, + "train_runtime": 64.7823, + "train_samples_per_second": 6.175, + "train_steps_per_second": 1.544 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2191347477905408.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f390eebadead91a612deb884067d1c4eb03c8131 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:051e6a1f5ba70e527bca65a2cf009359f5ce3064207e72a4ada270123beb41d9 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..2df3b5cae3920230980719f2fd4c308609f88bbc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3f576e7e654caeec4f265c632498460af8c22c262eb83d10734219f0186dbfaf +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..16e366b3bc3b9b9efb0d7caa827088a91f43d1ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:971321fe4c434b3e45168f2b1ec956068a3b0eeef06d78b4d7b095faaea7c6f4 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..c10e1436c46fa4cb0012220d31222a811270a768 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b09d5067ff7566b044ec865335563854b643ad29a4e43e158dc25473cfaab998 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..8196c50e942811c6bb72c0ee63418582692d1386 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a1ddf4094d69b5a5796c5559c6479d1bf03649b6debf99b0a9b8545943d1962a +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..7d857cff2533b02df39208e05b79ed42543fd0cb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a345facc12e980845071c2b28a1c0daeedf90e7983145dd31c029949febc5a9d +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..e343e1b5664e5e12bec257977a9580d45e3daf0c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:15f19b894addde6e68951ef07ab9d789a5a0dc973be735a34f8dac0e2aed0740 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..ed76d497de41253b0ce72ee5737dffb302197e5d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:49077774a25e071b758cb742c3deecb74c13d81c70dbc5ba613cf0335cc7e501 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..383cd8e92c287d8760b795edcd54c5b3a76ee432 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/12_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.0440428256988525, + "learning_rate": 2e-05, + "loss": 0.0319, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.003649517660960555, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.07765121757984161, + "learning_rate": 2e-05, + "loss": 0.0136, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.6474246978759766, + "learning_rate": 2e-05, + "loss": 0.0233, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.5474910736083984, + "learning_rate": 2e-05, + "loss": 0.0248, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.09228961914777756, + "learning_rate": 2e-05, + "loss": 0.0032, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.857515811920166, + "learning_rate": 2e-05, + "loss": 0.5061, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.1475517749786377, + "learning_rate": 2e-05, + "loss": 0.2934, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.27937111258506775, + "learning_rate": 2e-05, + "loss": 0.1483, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.492009162902832, + "learning_rate": 2e-05, + "loss": 0.2008, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.19078873097896576, + "learning_rate": 2e-05, + "loss": 0.1083, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.16366007924079895, + "learning_rate": 2e-05, + "loss": 0.011, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.4017130136489868, + "learning_rate": 2e-05, + "loss": 0.074, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.319075584411621, + "learning_rate": 2e-05, + "loss": 0.0607, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.022196315228939056, + "learning_rate": 2e-05, + "loss": 0.0073, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.07995142042636871, + "learning_rate": 2e-05, + "loss": 0.0089, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.12750348448753357, + "learning_rate": 2e-05, + "loss": 0.0301, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.20020344853401184, + "learning_rate": 2e-05, + "loss": 0.0414, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.27946731448173523, + "learning_rate": 2e-05, + "loss": 0.0191, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.713668704032898, + "learning_rate": 2e-05, + "loss": 0.0428, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.1871607303619385, + "learning_rate": 2e-05, + "loss": 0.2085, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.029789531603455544, + "learning_rate": 2e-05, + "loss": 0.0246, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.477757215499878, + "learning_rate": 2e-05, + "loss": 0.2057, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.3392833471298218, + "learning_rate": 2e-05, + "loss": 0.0527, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.05242856964468956, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.11674490571022034, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.05272448807954788, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.07284173369407654, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.016777602955698967, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.0905802994966507, + "learning_rate": 2e-05, + "loss": 0.0053, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.00765503803268075, + "learning_rate": 2e-05, + "loss": 0.0052, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.00952236633747816, + "learning_rate": 2e-05, + "loss": 0.005, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.08325570076704025, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.0055264234542847, + "learning_rate": 2e-05, + "loss": 0.0404, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.00400048540905118, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.07774416357278824, + "learning_rate": 2e-05, + "loss": 0.3443, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.07318832725286484, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.010225958190858364, + "learning_rate": 2e-05, + "loss": 0.0177, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.33187806606292725, + "learning_rate": 2e-05, + "loss": 0.245, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.0037391006480902433, + "learning_rate": 2e-05, + "loss": 0.0059, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.05806802958250046, + "learning_rate": 2e-05, + "loss": 0.0025, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.05544259026646614, + "learning_rate": 2e-05, + "loss": 0.0074, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.005225593224167824, + "learning_rate": 2e-05, + "loss": 0.0107, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.5481507778167725, + "learning_rate": 2e-05, + "loss": 0.1947, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.0036578518338501453, + "learning_rate": 2e-05, + "loss": 0.0094, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.14042843878269196, + "learning_rate": 2e-05, + "loss": 0.048, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.053851127624512, + "learning_rate": 2e-05, + "loss": 0.4281, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.03389669582247734, + "learning_rate": 2e-05, + "loss": 0.4295, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.07074581831693649, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.5740609169006348, + "learning_rate": 2e-05, + "loss": 0.0384, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5285302658662400.0, + "train_loss": 0.08000924110412598, + "train_runtime": 114.0475, + "train_samples_per_second": 3.507, + "train_steps_per_second": 0.877 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5285302658662400.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..1e582a5c369d5245fcbfb8bc5192da51dc20175b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:637378a39b5f003a03e125898e1002b87d4ab6d0636a9c60b6a2ccef5b19db03 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..c6e24f1b45d78872e59328d145bfb31b136cf32d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fbb3f8628f26b376dc3b4b202bb234092a57e5d75c0b54b704e89ae3366af7f +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..91b43ca0d44a8f3fa3441414369cecac0e5bae73 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aac144ea2fed7d29bcc77ea5db2a30fa34a72aad176a6bdb76ab676dba6617bb +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e76c25073328d04ca34afc4d04d550433d5d7eaf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8335a5c4bd4eb0cd9f4e07b45e9449081b22fa65b270335953580a115c6ece04 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..fe78e9f3a9bd4664a1f2e1e688e9995896b1689a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:84cce9ed1f003af88095b4136ea510f73336b73ea0e02227c50c42769d0f81a8 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d60a9fc1ff708c6360aa4f402678c8448945837d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:60630de6d10254a4873451d666da6fb598760e7758b39e557337e8d13bf3fe4b +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b7d123bc9654990f422da2d7823548ca9ee7a6ee --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:02f8977fe0686a13de1c9e847fa311541b200948da6170471df613c4cfca08a9 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..80e1de89d4f9ea8787fa5904ceec7dbee6474b72 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a54c6ea48ec7cc423607cd672a822256100d4663292366072542f8fc5a9cf525 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..ffc92b628e03aa2693ee489e2060e15e8e1a4a77 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/13_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.545411586761475, + "learning_rate": 2e-05, + "loss": 0.3849, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.6999170780181885, + "learning_rate": 2e-05, + "loss": 0.2808, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.1701340675354004, + "learning_rate": 2e-05, + "loss": 0.2401, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.10392165184021, + "learning_rate": 2e-05, + "loss": 0.0855, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.7305335998535156, + "learning_rate": 2e-05, + "loss": 0.075, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.30869001150131226, + "learning_rate": 2e-05, + "loss": 0.0358, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.3211586475372314, + "learning_rate": 2e-05, + "loss": 0.1096, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.4351699352264404, + "learning_rate": 2e-05, + "loss": 0.3176, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.4861433506011963, + "learning_rate": 2e-05, + "loss": 0.1712, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.4178895950317383, + "learning_rate": 2e-05, + "loss": 1.2003, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 6.215935230255127, + "learning_rate": 2e-05, + "loss": 0.6758, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 5.790246963500977, + "learning_rate": 2e-05, + "loss": 0.201, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.1957221031188965, + "learning_rate": 2e-05, + "loss": 0.8769, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.0322773456573486, + "learning_rate": 2e-05, + "loss": 0.107, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 4.61569881439209, + "learning_rate": 2e-05, + "loss": 0.3983, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.6823124885559082, + "learning_rate": 2e-05, + "loss": 0.1093, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.1920750141143799, + "learning_rate": 2e-05, + "loss": 0.2482, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.266702175140381, + "learning_rate": 2e-05, + "loss": 0.3447, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.700031042098999, + "learning_rate": 2e-05, + "loss": 0.1297, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.451580286026001, + "learning_rate": 2e-05, + "loss": 0.2859, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.04963172599673271, + "learning_rate": 2e-05, + "loss": 0.2405, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.6328516006469727, + "learning_rate": 2e-05, + "loss": 0.3162, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.03167015314102173, + "learning_rate": 2e-05, + "loss": 0.0109, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.2080798149108887, + "learning_rate": 2e-05, + "loss": 0.1503, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.22369244694709778, + "learning_rate": 2e-05, + "loss": 0.1677, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.292408466339111, + "learning_rate": 2e-05, + "loss": 0.3715, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.625317096710205, + "learning_rate": 2e-05, + "loss": 0.3785, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.4896113872528076, + "learning_rate": 2e-05, + "loss": 0.185, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.05531296133995056, + "learning_rate": 2e-05, + "loss": 0.0368, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.10100290179252625, + "learning_rate": 2e-05, + "loss": 0.0549, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.9116363525390625, + "learning_rate": 2e-05, + "loss": 0.1737, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.45428645610809326, + "learning_rate": 2e-05, + "loss": 0.089, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.323458433151245, + "learning_rate": 2e-05, + "loss": 0.4212, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.8958280086517334, + "learning_rate": 2e-05, + "loss": 0.2662, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.2603542804718018, + "learning_rate": 2e-05, + "loss": 0.2469, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.5217472910881042, + "learning_rate": 2e-05, + "loss": 0.0406, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.7623229622840881, + "learning_rate": 2e-05, + "loss": 0.6375, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.6854968070983887, + "learning_rate": 2e-05, + "loss": 0.0898, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.120392322540283, + "learning_rate": 2e-05, + "loss": 0.2088, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.7161530256271362, + "learning_rate": 2e-05, + "loss": 0.0939, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.7854943871498108, + "learning_rate": 2e-05, + "loss": 0.0957, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 6.270991325378418, + "learning_rate": 2e-05, + "loss": 0.8002, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 6.876936912536621, + "learning_rate": 2e-05, + "loss": 0.6731, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.7822352647781372, + "learning_rate": 2e-05, + "loss": 0.1663, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.8954455852508545, + "learning_rate": 2e-05, + "loss": 0.1198, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.265428304672241, + "learning_rate": 2e-05, + "loss": 0.5632, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.066023826599121, + "learning_rate": 2e-05, + "loss": 0.0228, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.620649576187134, + "learning_rate": 2e-05, + "loss": 0.2553, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.45771992206573486, + "learning_rate": 2e-05, + "loss": 0.3147, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.6073453426361084, + "learning_rate": 2e-05, + "loss": 0.272, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5322110830379008.0, + "train_loss": 0.27480587005615237, + "train_runtime": 114.0241, + "train_samples_per_second": 3.508, + "train_steps_per_second": 0.877 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5322110830379008.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..779b6e7fc1b737d2f45deb55952257980765f981 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:354db91f8782dacf84ee146cd14cd138581feeb1a0c36ecae828968743657380 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..1f7c74651251ac310dd7ec27b0fc21ed1f421fcb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:88613296a37ce322fdeca72ee5b95870f625b09ea2f241d0c9e6f592e66683d5 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..7c45b43a3261bc2d27d5e35be7257c391ce920c7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72fb06b83a98570ee2cc5cfc9536344e92d6e3d4d9b30e0e83bf4efa8a4fb6ec +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..06a4bbbc74b179ea001a85e9067b3ed88de2fb36 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:264c9d7a952746e4a57e7eb0eefa0b08d6d51f9b368c9cca7caafeaf232f52e2 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..77580bf4eb6b79d0d2d338f2a0716f14672f6092 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6a38d3151fff276e75ed9f35692faaaa94a57b6fccb9085e2a4ae5ab7c0d1fa +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..fe1a62a0b31eacac6c7ab9c62fcebf0aca6aeb26 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5822ca66329bec951cf5df5490475473f79dd43ddadd5df111f78df0cc621d66 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..f25891a16bb18202786a106dd0fad8469bc4dbe5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7a99c3a7b83964686d8a4b8fe853c7132a9cd59ff765a55a9ccd9c230bd1127f +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8b3636717035bb2b58f152fd1da33e2b7515520c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:db2162ba2c3ed720a1f573bef1a74a7a8b9fe27e2ff3e3e67bd0d5e5dd12cf27 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..274436ad1b6fcce247079c0226156fbcb8a8602e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/14_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.7358377575874329, + "learning_rate": 2e-05, + "loss": 0.422, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.581648111343384, + "learning_rate": 2e-05, + "loss": 0.3973, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.20724786818027496, + "learning_rate": 2e-05, + "loss": 0.1099, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.04636262357234955, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.7850165367126465, + "learning_rate": 2e-05, + "loss": 0.7581, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.1152070760726929, + "learning_rate": 2e-05, + "loss": 0.24, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.826275587081909, + "learning_rate": 2e-05, + "loss": 0.3553, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.024863580241799355, + "learning_rate": 2e-05, + "loss": 0.0052, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.490124464035034, + "learning_rate": 2e-05, + "loss": 0.0926, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.4378279149532318, + "learning_rate": 2e-05, + "loss": 0.0295, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.23534899950027466, + "learning_rate": 2e-05, + "loss": 0.0295, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.03806605935096741, + "learning_rate": 2e-05, + "loss": 0.0021, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 6.515997886657715, + "learning_rate": 2e-05, + "loss": 0.6703, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 6.961805820465088, + "learning_rate": 2e-05, + "loss": 0.3221, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.7749924659729004, + "learning_rate": 2e-05, + "loss": 0.2122, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.8729695677757263, + "learning_rate": 2e-05, + "loss": 0.0729, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.16095773875713348, + "learning_rate": 2e-05, + "loss": 0.139, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.5947050452232361, + "learning_rate": 2e-05, + "loss": 0.0363, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.13164369761943817, + "learning_rate": 2e-05, + "loss": 0.0145, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.7745718955993652, + "learning_rate": 2e-05, + "loss": 0.3666, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.3166494071483612, + "learning_rate": 2e-05, + "loss": 0.1203, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.4466262757778168, + "learning_rate": 2e-05, + "loss": 0.445, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.3789776861667633, + "learning_rate": 2e-05, + "loss": 0.034, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.0951488018035889, + "learning_rate": 2e-05, + "loss": 0.0579, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.02851537987589836, + "learning_rate": 2e-05, + "loss": 0.1831, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.035258397459983826, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.8462300300598145, + "learning_rate": 2e-05, + "loss": 0.0646, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.648222804069519, + "learning_rate": 2e-05, + "loss": 0.0782, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.3166184425354004, + "learning_rate": 2e-05, + "loss": 0.2181, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.014277494512498379, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7237198948860168, + "learning_rate": 2e-05, + "loss": 0.2877, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.42094871401786804, + "learning_rate": 2e-05, + "loss": 0.0246, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.2924734354019165, + "learning_rate": 2e-05, + "loss": 0.0633, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.9884778261184692, + "learning_rate": 2e-05, + "loss": 0.1732, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.731753945350647, + "learning_rate": 2e-05, + "loss": 0.2164, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.19647128880023956, + "learning_rate": 2e-05, + "loss": 0.0205, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.6730728149414062, + "learning_rate": 2e-05, + "loss": 0.1322, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.702977418899536, + "learning_rate": 2e-05, + "loss": 0.1588, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.6952839493751526, + "learning_rate": 2e-05, + "loss": 0.1075, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.721588373184204, + "learning_rate": 2e-05, + "loss": 0.0525, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.5381120443344116, + "learning_rate": 2e-05, + "loss": 0.0483, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.8108348846435547, + "learning_rate": 2e-05, + "loss": 0.0544, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.5784426331520081, + "learning_rate": 2e-05, + "loss": 0.0195, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.1764838546514511, + "learning_rate": 2e-05, + "loss": 0.0208, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.6076014041900635, + "learning_rate": 2e-05, + "loss": 0.0998, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.9815555810928345, + "learning_rate": 2e-05, + "loss": 0.0241, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.1499319076538086, + "learning_rate": 2e-05, + "loss": 0.0157, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.2355562448501587, + "learning_rate": 2e-05, + "loss": 0.026, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.4025782644748688, + "learning_rate": 2e-05, + "loss": 0.0109, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.0052330149337649345, + "learning_rate": 2e-05, + "loss": 0.0345, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5337401710870528.0, + "train_loss": 0.14149075031280517, + "train_runtime": 120.4797, + "train_samples_per_second": 3.32, + "train_steps_per_second": 0.83 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5337401710870528.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f0c1945fd216aa396b5b610bc86a54a7b0f93654 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4537a68e42b08ffd746110cbd11b33db9f7b96e3aee684a086fbfebef92de694 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..43cf969d70159d4df73ae65e7dcf2d1ea2dfaeef --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09499662bbe0d377de20a8ca27a2bdc6c1f27fd8c09a1391d26fda52a2661939 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..06e178fe0fccaa7d371c02193686d15d87e6686a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:184276c5d5fb3fd45dfe622314e6d1b830285b1ec92fc018decb9e1dff379629 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..8671cac42451731ecf5f26991e44775791c57eaa --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd643757fa412056c23eabe7e1da5a3df47796564d5062f6c1163459ebd73d9d +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..95a929ffc5ac57dd122e4792a8641bb7e84f10da --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafcf636e19eb7d5acfb633345efd2c20aeec4a163e0418380f6008686b2e54a +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..91162a0daed5be6ee46f42e769f0c70ebc06cdd9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d2424a0ca3a9248409218336e6af0718634093d36cfd210064d54087eb41e5d8 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..90296a46f9202c9aca76b954308db70c8ccdf180 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5f26b57f79fbde269125357ed1c7f0fc4dd67e591cf57ea5f0d55558ba70fde3 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..9214186dcb7922b4dce38e35699f77e99dba74a1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:50e8acf5896837685caa936fd27d061be00dd8e068e8a31e1cc803d30aad97eb +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..736ad383fe06b5c1a1342df94e065444e9206dd2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/15_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.4749128818511963, + "learning_rate": 2e-05, + "loss": 0.1664, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 8.862194061279297, + "learning_rate": 2e-05, + "loss": 0.5214, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 5.990468502044678, + "learning_rate": 2e-05, + "loss": 0.4672, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.7150177955627441, + "learning_rate": 2e-05, + "loss": 0.1249, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8785606622695923, + "learning_rate": 2e-05, + "loss": 0.0543, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.0889270305633545, + "learning_rate": 2e-05, + "loss": 0.2588, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.7522593140602112, + "learning_rate": 2e-05, + "loss": 0.03, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.42291760444641113, + "learning_rate": 2e-05, + "loss": 0.389, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.3472583293914795, + "learning_rate": 2e-05, + "loss": 0.0876, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.589219331741333, + "learning_rate": 2e-05, + "loss": 0.1295, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 5.906162738800049, + "learning_rate": 2e-05, + "loss": 0.4197, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 4.201356887817383, + "learning_rate": 2e-05, + "loss": 0.2237, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.554701805114746, + "learning_rate": 2e-05, + "loss": 0.4366, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.5249311923980713, + "learning_rate": 2e-05, + "loss": 0.2653, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 4.1731719970703125, + "learning_rate": 2e-05, + "loss": 0.2499, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.5751137733459473, + "learning_rate": 2e-05, + "loss": 0.0718, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.005580425262451, + "learning_rate": 2e-05, + "loss": 0.4738, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.2529420852661133, + "learning_rate": 2e-05, + "loss": 0.0712, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.1782426834106445, + "learning_rate": 2e-05, + "loss": 0.3885, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.842854022979736, + "learning_rate": 2e-05, + "loss": 0.2207, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.4326973259449005, + "learning_rate": 2e-05, + "loss": 0.0475, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.397282123565674, + "learning_rate": 2e-05, + "loss": 0.5645, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 6.812735557556152, + "learning_rate": 2e-05, + "loss": 0.1161, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.549821138381958, + "learning_rate": 2e-05, + "loss": 0.0422, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 5.42340087890625, + "learning_rate": 2e-05, + "loss": 0.7025, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.5064210891723633, + "learning_rate": 2e-05, + "loss": 0.282, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.1591339260339737, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.6468251943588257, + "learning_rate": 2e-05, + "loss": 0.0215, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.267123222351074, + "learning_rate": 2e-05, + "loss": 0.3784, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8600457906723022, + "learning_rate": 2e-05, + "loss": 0.0937, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 10.030192375183105, + "learning_rate": 2e-05, + "loss": 0.4602, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.5103938579559326, + "learning_rate": 2e-05, + "loss": 0.0981, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.9807940721511841, + "learning_rate": 2e-05, + "loss": 0.0865, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 8.568153381347656, + "learning_rate": 2e-05, + "loss": 0.5535, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.5489164590835571, + "learning_rate": 2e-05, + "loss": 0.5308, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 3.896538734436035, + "learning_rate": 2e-05, + "loss": 0.3199, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.36668315529823303, + "learning_rate": 2e-05, + "loss": 0.1478, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 4.95508337020874, + "learning_rate": 2e-05, + "loss": 0.3202, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.1562553644180298, + "learning_rate": 2e-05, + "loss": 0.0468, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.8538686633110046, + "learning_rate": 2e-05, + "loss": 0.1461, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.3485958576202393, + "learning_rate": 2e-05, + "loss": 0.1832, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.334911823272705, + "learning_rate": 2e-05, + "loss": 0.1918, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.204462051391602, + "learning_rate": 2e-05, + "loss": 0.5045, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.736405372619629, + "learning_rate": 2e-05, + "loss": 0.7114, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.910045623779297, + "learning_rate": 2e-05, + "loss": 0.2643, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.5672945976257324, + "learning_rate": 2e-05, + "loss": 0.2352, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.8731167316436768, + "learning_rate": 2e-05, + "loss": 0.051, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.5115487575531006, + "learning_rate": 2e-05, + "loss": 0.8067, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.768833160400391, + "learning_rate": 2e-05, + "loss": 0.2662, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.9175615310668945, + "learning_rate": 2e-05, + "loss": 0.0816, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2213537237696512.0, + "train_loss": 0.26631745338439944, + "train_runtime": 75.1558, + "train_samples_per_second": 5.322, + "train_steps_per_second": 1.331 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2213537237696512.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..df62f06514b76572802c812602ea04ed69747b7b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c704a4187e7a0c012f14f82d84bbf12b0ea847f6971f547a03c730f560079b9 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ae01c2e7e2e800c2b34fa24320645ac55415de5a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:562c5e8096eb204a914a76a20171a6b4c35dd325c2a6f21b6204d43eb8f7727f +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..47ce85b3abceaa3785b3b5117fd98a09b69a3bc3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9a05d02814da8c818ebab2e96fc53b1b7befdbd59241b829977e94804238504e +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..cd2cbf9fd1ed053e32ea1d14a7d2c85224f7acf1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c5d6cedd2c0f44b7e251903ba7878ee4b29f109b07e21e36a366ac14dbbde9c1 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..c3b06975a65b3966d97540747eb46d39bd03409b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b891cbcfeadd1630d6879b44102080272a4eac529424d80358440b2483f8bf52 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..7e5dfae2c8cc744bf859340b732e0fb7d5b57b8a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:991a13b1c91cb425fea2eac6f70cec289c3a5ff61c4715e3c951e854d6c0809c +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..f5e3e622a2a5c16872f25ef6d70b93e2d09ca4d1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f2dc79c6404c0c996476f33ea9b8e14bf8c0ba3738dc75419b5dec3ebdaa1af4 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..25ff61b74c2efc85cf75950a7013a169fdab37dc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:30c59408ad1099cb43f73191a326330fe56f5719a754ab8bc381e9fc67fef0d7 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..743a05067a3d0b21ee332773a17dc51dbd829ecb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/16_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.5247721672058105, + "learning_rate": 2e-05, + "loss": 0.2837, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.3108275830745697, + "learning_rate": 2e-05, + "loss": 0.0301, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 13.340131759643555, + "learning_rate": 2e-05, + "loss": 0.8614, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 4.815690040588379, + "learning_rate": 2e-05, + "loss": 0.2022, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.714431285858154, + "learning_rate": 2e-05, + "loss": 0.2346, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.055443789809942245, + "learning_rate": 2e-05, + "loss": 0.0455, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.043651819229126, + "learning_rate": 2e-05, + "loss": 0.1245, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.9224467277526855, + "learning_rate": 2e-05, + "loss": 0.3654, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.758558988571167, + "learning_rate": 2e-05, + "loss": 0.0603, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.22429201006889343, + "learning_rate": 2e-05, + "loss": 0.0258, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 5.3818230628967285, + "learning_rate": 2e-05, + "loss": 0.6414, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.30497607588768005, + "learning_rate": 2e-05, + "loss": 0.0908, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.457232475280762, + "learning_rate": 2e-05, + "loss": 0.5499, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.3166507482528687, + "learning_rate": 2e-05, + "loss": 0.1442, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.3585424423217773, + "learning_rate": 2e-05, + "loss": 0.4432, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 6.761358261108398, + "learning_rate": 2e-05, + "loss": 0.425, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.2732912600040436, + "learning_rate": 2e-05, + "loss": 0.2753, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.5978434085845947, + "learning_rate": 2e-05, + "loss": 0.2123, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.907473087310791, + "learning_rate": 2e-05, + "loss": 0.0977, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.6676077842712402, + "learning_rate": 2e-05, + "loss": 0.257, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.21937216818332672, + "learning_rate": 2e-05, + "loss": 0.0137, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.152346134185791, + "learning_rate": 2e-05, + "loss": 0.1916, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.4615230560302734, + "learning_rate": 2e-05, + "loss": 0.0779, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.3725461959838867, + "learning_rate": 2e-05, + "loss": 0.1205, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.16937066614627838, + "learning_rate": 2e-05, + "loss": 0.0795, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.8615169525146484, + "learning_rate": 2e-05, + "loss": 0.0433, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.7729310989379883, + "learning_rate": 2e-05, + "loss": 0.0697, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 6.881181716918945, + "learning_rate": 2e-05, + "loss": 0.4339, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.623704433441162, + "learning_rate": 2e-05, + "loss": 0.2259, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.2712254524230957, + "learning_rate": 2e-05, + "loss": 0.0775, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.8995764851570129, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.920132875442505, + "learning_rate": 2e-05, + "loss": 0.5304, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.4464337825775146, + "learning_rate": 2e-05, + "loss": 0.0938, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 12.049845695495605, + "learning_rate": 2e-05, + "loss": 1.5502, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.604862928390503, + "learning_rate": 2e-05, + "loss": 0.227, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 8.990845680236816, + "learning_rate": 2e-05, + "loss": 0.4252, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 6.3416032791137695, + "learning_rate": 2e-05, + "loss": 0.7439, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.05386652797460556, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.0817679166793823, + "learning_rate": 2e-05, + "loss": 0.0196, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.22007165849208832, + "learning_rate": 2e-05, + "loss": 0.0293, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 4.586897850036621, + "learning_rate": 2e-05, + "loss": 0.5778, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.766180694103241, + "learning_rate": 2e-05, + "loss": 0.2481, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.1154505014419556, + "learning_rate": 2e-05, + "loss": 0.0267, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 4.922826290130615, + "learning_rate": 2e-05, + "loss": 0.248, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.8402912616729736, + "learning_rate": 2e-05, + "loss": 0.4094, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.000579833984375, + "learning_rate": 2e-05, + "loss": 0.1154, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 8.067148208618164, + "learning_rate": 2e-05, + "loss": 0.8396, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 5.565702438354492, + "learning_rate": 2e-05, + "loss": 0.1246, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.5809526443481445, + "learning_rate": 2e-05, + "loss": 0.0977, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.8673484325408936, + "learning_rate": 2e-05, + "loss": 0.1246, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206762417520640.0, + "train_loss": 0.2633969497680664, + "train_runtime": 65.3498, + "train_samples_per_second": 6.121, + "train_steps_per_second": 1.53 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206762417520640.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..89211795bf0aedf5f7bbd023ce0b96423097c41a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3fc20d582960602a0db2baecf6e24341a16b4d6a2ccf50fe6147f005864c2784 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..265fa2f6aaef1c8f17982ff826d6e05b2b677470 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:defe11902846d18d40a74168cb71854343a822115a7898420818aa523e68977b +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..fef08f0491719fb3a7bc87f8d20793aa5f3a3451 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:decfd7da58a9f3582d079d520f30baa4c5d90b706c05cf95e16319818065a96e +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..a187235c2ec31689af20c0e5939453ba49f6e4ae --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:42cde2160f57197cd446304e672c26812bb6613df9d78216ed9dc0f32c1cc727 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..6df7d5e30e69bc6f9173a46714d8756e7dfbdaa8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:78961e064ccf7853afd8ba3d80732f54435264cf010125ee38a4b5b16464299f +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..8dea6b6990157cdc0e6ca10a610ee64720042727 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aba0032dc85997b78cd7257b53fe0abdafd0b8024fb99453463ca4a1309215ef +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c68a56b6ecb9098ca6f29e44d8b3a990bc55f8d2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3875fde3c1272e6f70c3e405bb60beab31c79cf1906c901d189f324bc243b720 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..20e3a77b7c1d912eb917439cfdb3440149cbb585 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cac8c2b2f491f8364c5557f92265c33df0a894c578b031c0fa6fd777235caa9d +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..f5bdce952a5a06c3366caac49f9dd97c916a3ab9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/17_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.9309191703796387, + "learning_rate": 2e-05, + "loss": 0.0483, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.144968509674072, + "learning_rate": 2e-05, + "loss": 0.4469, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 3.9196979999542236, + "learning_rate": 2e-05, + "loss": 0.1995, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.1354676485061646, + "learning_rate": 2e-05, + "loss": 0.0726, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.146620750427246, + "learning_rate": 2e-05, + "loss": 0.1982, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.015084743499756, + "learning_rate": 2e-05, + "loss": 0.2976, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.182155728340149, + "learning_rate": 2e-05, + "loss": 0.0361, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.102903366088867, + "learning_rate": 2e-05, + "loss": 0.1763, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 10.29223918914795, + "learning_rate": 2e-05, + "loss": 0.5212, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 8.495118141174316, + "learning_rate": 2e-05, + "loss": 0.751, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.731539487838745, + "learning_rate": 2e-05, + "loss": 0.1394, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 8.644320487976074, + "learning_rate": 2e-05, + "loss": 0.7443, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.8600324988365173, + "learning_rate": 2e-05, + "loss": 0.1115, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.18056751787662506, + "learning_rate": 2e-05, + "loss": 0.0091, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.9712249040603638, + "learning_rate": 2e-05, + "loss": 0.1475, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.323077440261841, + "learning_rate": 2e-05, + "loss": 0.0984, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.8479629755020142, + "learning_rate": 2e-05, + "loss": 0.0765, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.6739343404769897, + "learning_rate": 2e-05, + "loss": 0.0239, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.1887664794921875, + "learning_rate": 2e-05, + "loss": 0.2105, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.001166820526123, + "learning_rate": 2e-05, + "loss": 0.3232, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.5196139812469482, + "learning_rate": 2e-05, + "loss": 0.0268, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.9610042572021484, + "learning_rate": 2e-05, + "loss": 0.3582, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.168078899383545, + "learning_rate": 2e-05, + "loss": 0.0435, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 7.561069965362549, + "learning_rate": 2e-05, + "loss": 0.3503, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.111419200897217, + "learning_rate": 2e-05, + "loss": 0.3194, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.228247880935669, + "learning_rate": 2e-05, + "loss": 0.2119, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.595090866088867, + "learning_rate": 2e-05, + "loss": 0.2414, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 4.742661476135254, + "learning_rate": 2e-05, + "loss": 0.2564, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 7.1881537437438965, + "learning_rate": 2e-05, + "loss": 0.4555, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.82837438583374, + "learning_rate": 2e-05, + "loss": 0.4872, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 10.767409324645996, + "learning_rate": 2e-05, + "loss": 1.1565, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 6.279138088226318, + "learning_rate": 2e-05, + "loss": 0.978, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.4786481261253357, + "learning_rate": 2e-05, + "loss": 0.0693, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 7.9767746925354, + "learning_rate": 2e-05, + "loss": 0.3807, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.5030804872512817, + "learning_rate": 2e-05, + "loss": 0.0296, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 4.477349758148193, + "learning_rate": 2e-05, + "loss": 0.4666, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.0608134269714355, + "learning_rate": 2e-05, + "loss": 0.3884, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.140836715698242, + "learning_rate": 2e-05, + "loss": 0.7536, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 8.507833480834961, + "learning_rate": 2e-05, + "loss": 0.4483, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.081936001777649, + "learning_rate": 2e-05, + "loss": 0.1201, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.6277778148651123, + "learning_rate": 2e-05, + "loss": 0.1755, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 4.031922817230225, + "learning_rate": 2e-05, + "loss": 0.3693, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 3.7736613750457764, + "learning_rate": 2e-05, + "loss": 0.3403, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.2984522581100464, + "learning_rate": 2e-05, + "loss": 0.0386, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 5.806349754333496, + "learning_rate": 2e-05, + "loss": 0.4737, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.192534446716309, + "learning_rate": 2e-05, + "loss": 0.2145, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.18284006416797638, + "learning_rate": 2e-05, + "loss": 0.028, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.1052632331848145, + "learning_rate": 2e-05, + "loss": 0.129, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.020810067653656006, + "learning_rate": 2e-05, + "loss": 0.0596, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.6067790985107422, + "learning_rate": 2e-05, + "loss": 0.0758, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2211830323740672.0, + "train_loss": 0.2815584850311279, + "train_runtime": 71.4236, + "train_samples_per_second": 5.6, + "train_steps_per_second": 1.4 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2211830323740672.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5e66e98b14d8202a031201f6175be0165e41132d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b7bef768127b894b37e2dffbfd09cd26732e24eac8e882fb5682dcf7a58ff63b +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ee9100ebbe737bd16d715c22254c104334d5962f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:752436b73c53c133db8068b42299a64f7958ac34ec6daaa697a4632da1db5ccd +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..840f0aeef925ad0cf83303cecdaeac6295208f33 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:de896a8ef8ecb5a91c0aaab2a8ff019d21519147e81483f97fc7cab609ca4175 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..3e12c7f487c55821675b885fae7d5b7e19928e45 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:35a43071610a7aef17b8452b7d877afcd4e84453dad83136485936aaa0618449 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..09caac21769b203b30cf5918885cf03501e93800 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d477528feb4ca285190bd3f9f9102c143280c3251d7c1573761f669100c677e6 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b2208088d469b1f82c307017dc1e7f2bee4b2d7f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b5e23eb78ed44891b31dd1acc3c516a424f172377ded48c3cdfbc3bc4d653844 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..fdcdd39916e271f03d6fc0f119334d445ef26e33 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e443d668557b846262927ab80e0f37f37ae3e1b23afe87a9c3b5586a082583d1 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..36d63b49a12ff5039276adcf0a79ea9a76b310f9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a92a419d78fd8a90a0644a04806d7bad83e2489d331e06ecd2da6364f4a899c2 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..76b37b1c232b4aff7a2e62c4389bb032a1f512ef --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/18_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.3648357093334198, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.502422571182251, + "learning_rate": 2e-05, + "loss": 0.0816, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.8068466782569885, + "learning_rate": 2e-05, + "loss": 0.0232, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.07025347650051117, + "learning_rate": 2e-05, + "loss": 0.0026, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.056873202323913574, + "learning_rate": 2e-05, + "loss": 0.4384, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 5.214589595794678, + "learning_rate": 2e-05, + "loss": 0.102, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.8396440744400024, + "learning_rate": 2e-05, + "loss": 0.1251, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.059538722038269, + "learning_rate": 2e-05, + "loss": 0.1046, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.2488890886306763, + "learning_rate": 2e-05, + "loss": 0.0463, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.2230424880981445, + "learning_rate": 2e-05, + "loss": 0.141, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.03979168459773064, + "learning_rate": 2e-05, + "loss": 0.1476, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.40982499718666077, + "learning_rate": 2e-05, + "loss": 0.064, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.7363425493240356, + "learning_rate": 2e-05, + "loss": 0.0825, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.9172003269195557, + "learning_rate": 2e-05, + "loss": 0.2518, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.22544029355049133, + "learning_rate": 2e-05, + "loss": 0.0888, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 5.169614315032959, + "learning_rate": 2e-05, + "loss": 0.1464, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.6949363350868225, + "learning_rate": 2e-05, + "loss": 0.0355, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.30332151055336, + "learning_rate": 2e-05, + "loss": 0.1702, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.1602215766906738, + "learning_rate": 2e-05, + "loss": 0.0648, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.49905925989151, + "learning_rate": 2e-05, + "loss": 0.0388, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.1324706077575684, + "learning_rate": 2e-05, + "loss": 0.4652, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.046251799911260605, + "learning_rate": 2e-05, + "loss": 0.0056, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.491790533065796, + "learning_rate": 2e-05, + "loss": 0.2817, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.152531147003174, + "learning_rate": 2e-05, + "loss": 0.0716, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.661505937576294, + "learning_rate": 2e-05, + "loss": 0.0851, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.8545902371406555, + "learning_rate": 2e-05, + "loss": 0.0983, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 5.924610614776611, + "learning_rate": 2e-05, + "loss": 0.2678, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.16585147380828857, + "learning_rate": 2e-05, + "loss": 0.2567, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.3850327432155609, + "learning_rate": 2e-05, + "loss": 0.0947, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8235000967979431, + "learning_rate": 2e-05, + "loss": 0.0782, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.010295404121279716, + "learning_rate": 2e-05, + "loss": 0.0138, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.0673210620880127, + "learning_rate": 2e-05, + "loss": 0.0848, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.012797627598047256, + "learning_rate": 2e-05, + "loss": 0.0917, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.2304740697145462, + "learning_rate": 2e-05, + "loss": 0.0806, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 8.492612838745117, + "learning_rate": 2e-05, + "loss": 0.4658, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.07390482723712921, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.787769079208374, + "learning_rate": 2e-05, + "loss": 0.0529, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 10.277799606323242, + "learning_rate": 2e-05, + "loss": 0.3465, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 6.684095859527588, + "learning_rate": 2e-05, + "loss": 0.4043, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.5055056810379028, + "learning_rate": 2e-05, + "loss": 0.0125, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 5.695315361022949, + "learning_rate": 2e-05, + "loss": 0.2941, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.15820398926734924, + "learning_rate": 2e-05, + "loss": 0.2397, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.2003238201141357, + "learning_rate": 2e-05, + "loss": 0.0219, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.07542400807142258, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.5835587978363037, + "learning_rate": 2e-05, + "loss": 0.178, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.4093284606933594, + "learning_rate": 2e-05, + "loss": 0.0901, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 7.452977180480957, + "learning_rate": 2e-05, + "loss": 0.3038, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 7.80583381652832, + "learning_rate": 2e-05, + "loss": 0.1375, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.024793971329927444, + "learning_rate": 2e-05, + "loss": 0.0041, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 6.005980491638184, + "learning_rate": 2e-05, + "loss": 0.2792, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206568850391040.0, + "train_loss": 0.13983344793319702, + "train_runtime": 64.8795, + "train_samples_per_second": 6.165, + "train_steps_per_second": 1.541 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206568850391040.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..eb38a21d6f97f215d6cdc914389e8be8823bf194 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f40a36617c5e0c7016c40e4f2384d743b978481611f33d66322c7eafb0a83b5d +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..edef0f9d793c07c44124a110b036dac421fa73a9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b98bfe74c297e1761aa986296259170adc53f8f4a1e7affc763b30aea68e185 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..f0d124c8190817b29e042a014f24ada774b8da06 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7b956a03a4c777ef787dd9c5d8cda435f925404aef80ed1aa5d074dd80289c52 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..adf48d74aaa8ef11dbc6cdd74154d998db63ccc7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5348f638bf37175e8e84ef836b87698fc5759b8a653702a354ed923282ca7329 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..b916187b0cc5ffef9174952919c0ccd5bd44ba63 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4ca8cbc6fca2f550f1bd45f8e04d749c8caf49addf2efee68df6c5705d0bcd5b +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..138b1ad97f69319ce35aea9504e42d320f9956ce --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2626e07938d79f640578f2cf9b134ca1ada16078d5d82280b24530a8daea6a79 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ae770fa4387935623f45acae4f40fed8a51381d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:138439f6c6e23d70d094b65a53c4d3c2294dff89061ac5cde4d0ef236c13a5c4 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..1131d20c8b880819b7f315004f8bbd62005929b4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:82092804994a9d6e5ef29a73f9ea2437940c3df6fc35397b0aafe98d0cb83bb8 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..8b0cf737b36d04f6778198376c9ead562c041b0e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/19_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.3911495804786682, + "learning_rate": 2e-05, + "loss": 0.0844, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.3757598400115967, + "learning_rate": 2e-05, + "loss": 0.0843, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.7317628264427185, + "learning_rate": 2e-05, + "loss": 0.049, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.4361867904663086, + "learning_rate": 2e-05, + "loss": 0.1739, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.361838340759277, + "learning_rate": 2e-05, + "loss": 0.3539, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.3120211362838745, + "learning_rate": 2e-05, + "loss": 0.0471, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.1116899251937866, + "learning_rate": 2e-05, + "loss": 0.0729, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.6781142354011536, + "learning_rate": 2e-05, + "loss": 0.1386, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.23457714915275574, + "learning_rate": 2e-05, + "loss": 0.2478, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.2630733549594879, + "learning_rate": 2e-05, + "loss": 0.0154, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.9440138339996338, + "learning_rate": 2e-05, + "loss": 0.0767, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.8345701694488525, + "learning_rate": 2e-05, + "loss": 0.0765, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.309408664703369, + "learning_rate": 2e-05, + "loss": 0.1483, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.680891752243042, + "learning_rate": 2e-05, + "loss": 0.1144, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.831723928451538, + "learning_rate": 2e-05, + "loss": 0.1951, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.4016166627407074, + "learning_rate": 2e-05, + "loss": 0.5332, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.3797646760940552, + "learning_rate": 2e-05, + "loss": 0.2171, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.4994640052318573, + "learning_rate": 2e-05, + "loss": 0.105, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2403087615966797, + "learning_rate": 2e-05, + "loss": 0.0372, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.11302483081817627, + "learning_rate": 2e-05, + "loss": 0.0645, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.8516426682472229, + "learning_rate": 2e-05, + "loss": 0.0445, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.8383413553237915, + "learning_rate": 2e-05, + "loss": 0.5944, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.5927720665931702, + "learning_rate": 2e-05, + "loss": 0.1308, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.22132036089897156, + "learning_rate": 2e-05, + "loss": 0.2954, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.33681434392929077, + "learning_rate": 2e-05, + "loss": 0.2876, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.6511216163635254, + "learning_rate": 2e-05, + "loss": 0.1776, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.33043742179870605, + "learning_rate": 2e-05, + "loss": 0.217, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.08662908524274826, + "learning_rate": 2e-05, + "loss": 0.0137, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.2382876873016357, + "learning_rate": 2e-05, + "loss": 0.16, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.7894643545150757, + "learning_rate": 2e-05, + "loss": 0.0861, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.07965625822544098, + "learning_rate": 2e-05, + "loss": 0.1519, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.030141830444336, + "learning_rate": 2e-05, + "loss": 0.2856, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.878720760345459, + "learning_rate": 2e-05, + "loss": 0.252, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.0669201910495758, + "learning_rate": 2e-05, + "loss": 0.0915, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.927300214767456, + "learning_rate": 2e-05, + "loss": 0.4136, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.6932370066642761, + "learning_rate": 2e-05, + "loss": 0.1704, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.283788800239563, + "learning_rate": 2e-05, + "loss": 0.0247, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 6.912846565246582, + "learning_rate": 2e-05, + "loss": 0.5529, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.464448928833008, + "learning_rate": 2e-05, + "loss": 0.1552, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.000772714614868, + "learning_rate": 2e-05, + "loss": 0.1443, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.2344608306884766, + "learning_rate": 2e-05, + "loss": 0.3499, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.18237122893333435, + "learning_rate": 2e-05, + "loss": 0.0192, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0033660621847957373, + "learning_rate": 2e-05, + "loss": 0.0417, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.22952237725257874, + "learning_rate": 2e-05, + "loss": 0.3053, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3584758043289185, + "learning_rate": 2e-05, + "loss": 0.0528, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.045238733291626, + "learning_rate": 2e-05, + "loss": 0.0706, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.6022082567214966, + "learning_rate": 2e-05, + "loss": 0.1086, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.316227436065674, + "learning_rate": 2e-05, + "loss": 0.17, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.3068380355834961, + "learning_rate": 2e-05, + "loss": 0.0347, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.08089393377304077, + "learning_rate": 2e-05, + "loss": 0.3054, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293519879012352.0, + "train_loss": 0.1708518409729004, + "train_runtime": 111.7177, + "train_samples_per_second": 3.58, + "train_steps_per_second": 0.895 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293519879012352.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..07a8ad18465f03fb45b800182ae88a6a9109e149 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:07e831e1be8a3de5136c3af0a4456c79a674892f728645ce0ce0b879faa4e9cf +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..55064b39440a488cb70271da80229ad221252e56 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bc53c250ada0098c8ba9550aa104cb4c699155f45f587f9037f9288bbdd8f2b0 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..8738dd4e8be03953e513e19025e2158471188712 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1bd93ce47a6c310a3153753bbf7de074f5375d311f5a66faba417fd8a7c558c9 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..0910baef7ae4a33eea94c376a44f19a80f41e4c0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:27c97fa1cfc5e1b373cecda5ebfb978374a258188992d9db80b324ff5e661729 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5d95081066185c24e6705aed62b3374f226f998c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3eb6ffd8400e639b904ae85323a8c00d46032169f5dd450ff953d1c6b83f7ed6 +size 368442474 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..09bd60649c254d9d1a9a05a9c3e81cfe29478f07 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f83fe557d95bc2e40179adc3bbbf0edc9676b476b5df85da2c6241a6aa1d0bb7 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b17fb319e2407a6878e59231c586b477e40bd78b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2284729ba43be50088902af2c824ae03d8f7c36b79052b2f9980f83c54417a89 +size 368442474 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..e6197a98d2682863f104a7da2f52b2c845c575db --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5aafd60a920ed85dd5a5b2e007a70c46a16e9a373d9d4b8f52f10292d6166f22 +size 368442474 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..8743f3b75441f1cac52ebd4eb247e8a5070695f1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/1_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.033750955015420914, + "learning_rate": 2e-05, + "loss": 0.012, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.14480914175510406, + "learning_rate": 2e-05, + "loss": 0.0625, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.7304999828338623, + "learning_rate": 2e-05, + "loss": 0.0159, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.2602778673171997, + "learning_rate": 2e-05, + "loss": 0.0082, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.07092487066984177, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.48471954464912415, + "learning_rate": 2e-05, + "loss": 0.0363, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.05854671448469162, + "learning_rate": 2e-05, + "loss": 0.0014, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.20705780386924744, + "learning_rate": 2e-05, + "loss": 0.018, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 5.627336502075195, + "learning_rate": 2e-05, + "loss": 0.1969, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.01952861249446869, + "learning_rate": 2e-05, + "loss": 0.0012, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.042517539113759995, + "learning_rate": 2e-05, + "loss": 0.009, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.06011567637324333, + "learning_rate": 2e-05, + "loss": 0.0015, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.9543464183807373, + "learning_rate": 2e-05, + "loss": 0.1239, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.035090040415525436, + "learning_rate": 2e-05, + "loss": 0.0012, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.1751517504453659, + "learning_rate": 2e-05, + "loss": 0.0053, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.36874547600746155, + "learning_rate": 2e-05, + "loss": 0.0161, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.523126482963562, + "learning_rate": 2e-05, + "loss": 0.033, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.0029429879505187273, + "learning_rate": 2e-05, + "loss": 0.0002, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.06031635403633118, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.2437900304794312, + "learning_rate": 2e-05, + "loss": 0.0272, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.6812621355056763, + "learning_rate": 2e-05, + "loss": 0.0116, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 5.380383014678955, + "learning_rate": 2e-05, + "loss": 0.1961, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.009075570851564407, + "learning_rate": 2e-05, + "loss": 0.0009, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.03006071038544178, + "learning_rate": 2e-05, + "loss": 0.0012, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.011029444634914398, + "learning_rate": 2e-05, + "loss": 0.0015, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.04194938763976097, + "learning_rate": 2e-05, + "loss": 0.027, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.4486429691314697, + "learning_rate": 2e-05, + "loss": 0.0495, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.038574088364839554, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.010087166912853718, + "learning_rate": 2e-05, + "loss": 0.1895, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.008964954875409603, + "learning_rate": 2e-05, + "loss": 0.0165, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.11719933897256851, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.049465227872133255, + "learning_rate": 2e-05, + "loss": 0.0079, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.028385698795318604, + "learning_rate": 2e-05, + "loss": 0.0014, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.07456324994564056, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.017502592876553535, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.3877832889556885, + "learning_rate": 2e-05, + "loss": 0.0071, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.1450010538101196, + "learning_rate": 2e-05, + "loss": 0.0215, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.023379459977149963, + "learning_rate": 2e-05, + "loss": 0.0014, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.0175793170928955, + "learning_rate": 2e-05, + "loss": 0.0404, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.008190244436264038, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.04602311924099922, + "learning_rate": 2e-05, + "loss": 0.0045, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0046734618954360485, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.010905645787715912, + "learning_rate": 2e-05, + "loss": 0.1783, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.04598071798682213, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.2268616259098053, + "learning_rate": 2e-05, + "loss": 0.0047, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.2769477367401123, + "learning_rate": 2e-05, + "loss": 0.014, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.014070335775613785, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.1433114856481552, + "learning_rate": 2e-05, + "loss": 0.0029, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.007433066610246897, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.10523644834756851, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2217003838341120.0, + "train_loss": 0.027327795028686524, + "train_runtime": 83.8996, + "train_samples_per_second": 4.768, + "train_steps_per_second": 1.192 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2217003838341120.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..19efb9210383235b84f305e2a49f741eb56c8fb3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8b14dc4591d03d22e44fac47259df67815a0a47c3ae03803a952fc29dce0381 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..fc1b80ef5c0440a3e644e124550be9f6d3f4868a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:45b78e59933ebf9bebf475592f1ca8b23d934ec50d901bbfb4e719fe7c192d39 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..443e437fcbd8c0ea19995129ce4e1cfd0b3aea1b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c2f7cb2da83d9e0403d0f2fa01724fdf811979e3548c3f96093c592970b133a0 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e538b96e7e3efa7cd5f5c80d19a8221adbb524aa --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1e1b9dad9379d419073da1654e6477c1bda6ca127bc20c00f8b28a7dba167a5b +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..3d304a323ae052d761f4b3ed32dbd0b222332ac1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ac59db0b6a0f6ec732173b34310fe3c7a3e01e32c44a81bb545876f5179359c5 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6bf546b7ed9d9cfc721e2c4d89eed7e2b62ea998 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fcb9584a6ea02d0605d455fd0aff966ff4b2fb81281aee9e6c08b3ad249f4143 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..60f7671b049f1c7314b0ae3918a509e6191441e8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7eb61ce6cb9f953be90b137ae9c6eb006d4bd71fe5a5f436cb6511c5619b036e +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8acf7c8a706ac8e16a893235c2e07fe428a69b27 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5239abdcf2101ff80d8392c672a63d1fe7e6333090cac342a80217335ca58c65 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..a49362414756290e5d8314764fceba4033f06b4b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/20_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.10622059553861618, + "learning_rate": 2e-05, + "loss": 0.032, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.863730549812317, + "learning_rate": 2e-05, + "loss": 0.0808, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.052234172821045, + "learning_rate": 2e-05, + "loss": 0.0741, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.5702794194221497, + "learning_rate": 2e-05, + "loss": 0.0233, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.695217132568359, + "learning_rate": 2e-05, + "loss": 0.509, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.12372706830501556, + "learning_rate": 2e-05, + "loss": 0.0063, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.5843559503555298, + "learning_rate": 2e-05, + "loss": 0.0489, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.0554932355880737, + "learning_rate": 2e-05, + "loss": 0.0752, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 4.897989749908447, + "learning_rate": 2e-05, + "loss": 0.2562, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.07643141597509384, + "learning_rate": 2e-05, + "loss": 0.168, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.5603024959564209, + "learning_rate": 2e-05, + "loss": 0.0188, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.2541738450527191, + "learning_rate": 2e-05, + "loss": 0.0166, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.9034767150878906, + "learning_rate": 2e-05, + "loss": 0.2747, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.21665063500404358, + "learning_rate": 2e-05, + "loss": 0.015, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.3188706636428833, + "learning_rate": 2e-05, + "loss": 0.0376, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.152240514755249, + "learning_rate": 2e-05, + "loss": 0.2127, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.11979620158672333, + "learning_rate": 2e-05, + "loss": 0.0531, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3752935528755188, + "learning_rate": 2e-05, + "loss": 0.0354, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.3420855402946472, + "learning_rate": 2e-05, + "loss": 0.0814, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.8663727045059204, + "learning_rate": 2e-05, + "loss": 0.0753, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.38089561462402344, + "learning_rate": 2e-05, + "loss": 0.1598, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.7258450984954834, + "learning_rate": 2e-05, + "loss": 0.0978, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.0675911903381348, + "learning_rate": 2e-05, + "loss": 0.0902, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.1028257608413696, + "learning_rate": 2e-05, + "loss": 0.0416, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.7641342878341675, + "learning_rate": 2e-05, + "loss": 0.1491, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.4352625906467438, + "learning_rate": 2e-05, + "loss": 0.0121, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.29529789090156555, + "learning_rate": 2e-05, + "loss": 0.1192, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.06029689684510231, + "learning_rate": 2e-05, + "loss": 0.0237, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.0123901991173625, + "learning_rate": 2e-05, + "loss": 0.1697, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.10268533229827881, + "learning_rate": 2e-05, + "loss": 0.0669, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.916579008102417, + "learning_rate": 2e-05, + "loss": 0.2984, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 5.2938618659973145, + "learning_rate": 2e-05, + "loss": 0.5691, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 6.4340314865112305, + "learning_rate": 2e-05, + "loss": 0.8205, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 4.767757415771484, + "learning_rate": 2e-05, + "loss": 0.1837, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.9049027562141418, + "learning_rate": 2e-05, + "loss": 0.0538, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.08755947649478912, + "learning_rate": 2e-05, + "loss": 0.0396, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.14745764434337616, + "learning_rate": 2e-05, + "loss": 0.0111, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 6.647001266479492, + "learning_rate": 2e-05, + "loss": 0.5654, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.4596193730831146, + "learning_rate": 2e-05, + "loss": 0.0314, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.0224930252879858, + "learning_rate": 2e-05, + "loss": 0.019, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.15938426554203033, + "learning_rate": 2e-05, + "loss": 0.0891, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.1378124952316284, + "learning_rate": 2e-05, + "loss": 0.0963, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.21264223754405975, + "learning_rate": 2e-05, + "loss": 0.0128, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.20153546333313, + "learning_rate": 2e-05, + "loss": 0.2233, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.0017821788787842, + "learning_rate": 2e-05, + "loss": 0.0915, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.6205344200134277, + "learning_rate": 2e-05, + "loss": 0.2615, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.6587465405464172, + "learning_rate": 2e-05, + "loss": 0.0399, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.1720789521932602, + "learning_rate": 2e-05, + "loss": 0.0063, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.814685821533203, + "learning_rate": 2e-05, + "loss": 0.1212, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.2568967044353485, + "learning_rate": 2e-05, + "loss": 0.0112, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293985824243712.0, + "train_loss": 0.131396541595459, + "train_runtime": 114.6927, + "train_samples_per_second": 3.488, + "train_steps_per_second": 0.872 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293985824243712.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..991ab5d3303afe72aac7013b4536aa04eaf656a4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6686b34105af496b44de7c39233496d41cc1f1b501c8cdb58ac7d623f209ae35 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..a3b59f8c6a9d2ebbbfba6d05cf5cf4ce6ba0127e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d50482766c068677fd9ed2baaf0ff2276202b461bcee458b7b6342b8deb16d25 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..6384697bc0b7c80064e36d06b6f6d09ca0a5aa20 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:310690769536aae96cc2fb3995f4be5e7d4e10751512339b482f4d58d5d59155 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..1bc3217c44a0c3abd887a9a84f4b28ce737da0a1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:15c5aa657b454137e778835e95817f013bc4ce25ff390cef5c815852cd87a1f6 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..4542aca64e642ad825538d011e59009b6829790f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a4bdd5197e17f9c92b4289acc9de1766978e7fbca08c4bbca5e30292ebc17834 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..77a278e771456f9c72307b55c45d3042c047304c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29e7174216699d61d634a1516fffc046c65a9fb1b88b3c231a27ef11ce1428b2 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c69d85b08588f851afb4db0a83c2929c14818205 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5567554a4be60271d94576e627e5c6afe0ea09bd46cb0561cff13c481a57267d +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..4a3fb7db9d77be773485804dd9d8b1ad155a13ff --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b40d9b391829b3b29b82c973c2a1aa61a8aa0351d694ff8e31327f9624774b4 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3e53207cec1843b07aab0528fc1aa96335338383 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/2_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.2468481063842773, + "learning_rate": 2e-05, + "loss": 0.225, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.3193917274475098, + "learning_rate": 2e-05, + "loss": 0.6914, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.7411956787109375, + "learning_rate": 2e-05, + "loss": 0.3792, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.595564126968384, + "learning_rate": 2e-05, + "loss": 0.3853, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.889966368675232, + "learning_rate": 2e-05, + "loss": 0.1447, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 3.5301265716552734, + "learning_rate": 2e-05, + "loss": 0.5283, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.40425944328308105, + "learning_rate": 2e-05, + "loss": 0.1666, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.3258602619171143, + "learning_rate": 2e-05, + "loss": 0.3093, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.795256495475769, + "learning_rate": 2e-05, + "loss": 0.2727, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.0333045721054077, + "learning_rate": 2e-05, + "loss": 0.4653, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.34709981083869934, + "learning_rate": 2e-05, + "loss": 0.0723, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.751197576522827, + "learning_rate": 2e-05, + "loss": 0.2379, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.143589496612549, + "learning_rate": 2e-05, + "loss": 0.5242, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.7344441413879395, + "learning_rate": 2e-05, + "loss": 0.173, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.3542690873146057, + "learning_rate": 2e-05, + "loss": 0.0818, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.806969165802002, + "learning_rate": 2e-05, + "loss": 0.1891, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.127126455307007, + "learning_rate": 2e-05, + "loss": 0.3181, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.0388799905776978, + "learning_rate": 2e-05, + "loss": 0.0597, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.970386028289795, + "learning_rate": 2e-05, + "loss": 0.6151, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 5.099111557006836, + "learning_rate": 2e-05, + "loss": 0.3871, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.2818540334701538, + "learning_rate": 2e-05, + "loss": 0.3386, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7701959609985352, + "learning_rate": 2e-05, + "loss": 0.0876, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.9809255003929138, + "learning_rate": 2e-05, + "loss": 0.2244, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.3327561616897583, + "learning_rate": 2e-05, + "loss": 0.1447, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.5163413286209106, + "learning_rate": 2e-05, + "loss": 0.1318, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 4.389053821563721, + "learning_rate": 2e-05, + "loss": 0.3344, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.475992441177368, + "learning_rate": 2e-05, + "loss": 0.2766, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 3.774714708328247, + "learning_rate": 2e-05, + "loss": 0.2385, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.961244821548462, + "learning_rate": 2e-05, + "loss": 0.0863, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.17517104744911194, + "learning_rate": 2e-05, + "loss": 0.1372, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.213293194770813, + "learning_rate": 2e-05, + "loss": 0.312, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.775750756263733, + "learning_rate": 2e-05, + "loss": 0.1412, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.8578336238861084, + "learning_rate": 2e-05, + "loss": 0.1118, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.7798962593078613, + "learning_rate": 2e-05, + "loss": 0.212, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.354323387145996, + "learning_rate": 2e-05, + "loss": 0.3791, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.8846645951271057, + "learning_rate": 2e-05, + "loss": 0.5015, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.8081483840942383, + "learning_rate": 2e-05, + "loss": 0.0307, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.0817105770111084, + "learning_rate": 2e-05, + "loss": 0.1684, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.9011064767837524, + "learning_rate": 2e-05, + "loss": 0.0486, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 6.007953643798828, + "learning_rate": 2e-05, + "loss": 0.5659, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 7.340831756591797, + "learning_rate": 2e-05, + "loss": 1.7358, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.3310190737247467, + "learning_rate": 2e-05, + "loss": 0.0139, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.6133837699890137, + "learning_rate": 2e-05, + "loss": 0.3107, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.556053400039673, + "learning_rate": 2e-05, + "loss": 0.3075, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.721317768096924, + "learning_rate": 2e-05, + "loss": 0.8916, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.20518216490745544, + "learning_rate": 2e-05, + "loss": 0.1093, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.3494601249694824, + "learning_rate": 2e-05, + "loss": 0.2544, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.5662200450897217, + "learning_rate": 2e-05, + "loss": 0.1454, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.5390119552612305, + "learning_rate": 2e-05, + "loss": 0.6951, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.003822922706604, + "learning_rate": 2e-05, + "loss": 0.092, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5221258941693952.0, + "train_loss": 0.30506493091583253, + "train_runtime": 114.0561, + "train_samples_per_second": 3.507, + "train_steps_per_second": 0.877 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5221258941693952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5d3a5add89fbbc052f869e972aac712520e3a849 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ab6aacce2c574179eb9f933bd31cd5328a5182203d1faadc925feea1eeb71ef6 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..0a3c40d398dea68811ca76d9cfa55b5109c06d0a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:71e066694668a5de84efaabfeb65c25a303fbdaf38d46e24b0415a92c7ad33cf +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..7a081739e6cf1e886ed888dd1892c71b05331aa5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b43e86b0626b057412e1a4c1ae9a418558122f2bec359df4b9b92366af23d5e +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..fc5c7a7d485b8c55dbb65a1dd9c699e1fa67ad72 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:512b3c2689b76760d8d25f33ee067ce9082747f6dfdd4ca429fa9501ef86fe7c +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..93f389d0c6fad2c2c812870d0c4a9e1ed046033d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56fd216cf43ecd7098710b072d5fe19bd4fe4c3229afb0c0edf1f50e11af4850 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..70d46febdd408873db7f4601f8f8f2bfaa1bc6ae --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:65d4a33711f9f87e7f06b90bcbe1784d2a76916669cc3f2a2e7b6fb3ff8621a9 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..3fc0acc90d3652e832dbe755baedda9777f5d4bc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c797b8b431e1fd21ef83820cec77d089ed10430cd258deb4e6a669151bca2c8a +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..d5c3968af896d06d71c23d25e80327a600db3c1d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f27ad55201e418db156f2bb03e254a648cab5534bd93ace2a0c2cf5f1c665101 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..5e19fe777014ed989db0edd048410f8ab46645ab --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/3_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.6850311756134033, + "learning_rate": 2e-05, + "loss": 0.875, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.1689105033874512, + "learning_rate": 2e-05, + "loss": 0.4041, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.6469191312789917, + "learning_rate": 2e-05, + "loss": 0.2119, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.7608203887939453, + "learning_rate": 2e-05, + "loss": 0.687, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.928190231323242, + "learning_rate": 2e-05, + "loss": 0.554, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 4.438275337219238, + "learning_rate": 2e-05, + "loss": 0.8873, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.774273157119751, + "learning_rate": 2e-05, + "loss": 0.7549, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.75411057472229, + "learning_rate": 2e-05, + "loss": 0.5626, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.209951877593994, + "learning_rate": 2e-05, + "loss": 0.1818, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.0785526037216187, + "learning_rate": 2e-05, + "loss": 0.3579, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.417556047439575, + "learning_rate": 2e-05, + "loss": 0.2787, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.9813671112060547, + "learning_rate": 2e-05, + "loss": 0.52, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.4448651075363159, + "learning_rate": 2e-05, + "loss": 0.1977, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.4670426845550537, + "learning_rate": 2e-05, + "loss": 0.2681, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.6569111347198486, + "learning_rate": 2e-05, + "loss": 0.4022, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.930314779281616, + "learning_rate": 2e-05, + "loss": 0.4122, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.464197874069214, + "learning_rate": 2e-05, + "loss": 0.6721, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.0680365562438965, + "learning_rate": 2e-05, + "loss": 0.7461, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.5992873907089233, + "learning_rate": 2e-05, + "loss": 0.2202, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.5932765007019043, + "learning_rate": 2e-05, + "loss": 0.5332, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 4.405234336853027, + "learning_rate": 2e-05, + "loss": 0.6582, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 4.52316951751709, + "learning_rate": 2e-05, + "loss": 0.5667, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.2715246677398682, + "learning_rate": 2e-05, + "loss": 0.1999, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.912550151348114, + "learning_rate": 2e-05, + "loss": 0.0638, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 3.167592763900757, + "learning_rate": 2e-05, + "loss": 0.5571, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.822756767272949, + "learning_rate": 2e-05, + "loss": 0.6832, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.09021428972482681, + "learning_rate": 2e-05, + "loss": 0.191, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.337232232093811, + "learning_rate": 2e-05, + "loss": 0.2529, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.1108458042144775, + "learning_rate": 2e-05, + "loss": 0.3, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 3.4974570274353027, + "learning_rate": 2e-05, + "loss": 0.4891, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.1878390312194824, + "learning_rate": 2e-05, + "loss": 0.3532, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.4413163661956787, + "learning_rate": 2e-05, + "loss": 0.3154, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.416686058044434, + "learning_rate": 2e-05, + "loss": 0.5578, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 4.5052170753479, + "learning_rate": 2e-05, + "loss": 0.8879, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.2083427906036377, + "learning_rate": 2e-05, + "loss": 0.3838, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.4879305362701416, + "learning_rate": 2e-05, + "loss": 0.6826, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.635220766067505, + "learning_rate": 2e-05, + "loss": 0.6741, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.399205446243286, + "learning_rate": 2e-05, + "loss": 0.2724, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.8843960762023926, + "learning_rate": 2e-05, + "loss": 0.3496, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.2431797981262207, + "learning_rate": 2e-05, + "loss": 0.1678, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.0537939071655273, + "learning_rate": 2e-05, + "loss": 0.1093, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.627830743789673, + "learning_rate": 2e-05, + "loss": 0.2421, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.7740010023117065, + "learning_rate": 2e-05, + "loss": 0.6196, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.1690895557403564, + "learning_rate": 2e-05, + "loss": 0.2946, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.2248268127441406, + "learning_rate": 2e-05, + "loss": 0.4723, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.157418727874756, + "learning_rate": 2e-05, + "loss": 0.3457, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.081430196762085, + "learning_rate": 2e-05, + "loss": 0.2712, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.431945562362671, + "learning_rate": 2e-05, + "loss": 0.1935, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.8719439506530762, + "learning_rate": 2e-05, + "loss": 0.3362, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.1026384830474854, + "learning_rate": 2e-05, + "loss": 0.2327, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5408561442062336.0, + "train_loss": 0.4290246391296387, + "train_runtime": 114.8613, + "train_samples_per_second": 3.482, + "train_steps_per_second": 0.871 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5408561442062336.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..fd18d65732ac500451c7653f4c64b72794bfa4e1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ecd1141009a114d7b3e041f5a04f55a451c1fd765c1854bc1bc37175d59eae99 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..d3fb86addaddbdc05d74da0122ff78a1f830922b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8547634c0043afa609680d8de0c56a8a3a180ffc11c4b56bade2b2da6bb7ddd9 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..ef7458acfc53aa0e5cffe16f681ec4dcb485d8e3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:014d9caf726698af62680b1afb1a47eabc06c77c38e61a238f8197d6e9fc7c02 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..8018034826010a2aa36b6f5cd4ab06f6f744ca0c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d1afc62545af99fa5019b219af589956a79f2d369f09eb36e56f7af0960e3989 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..f76967fdc566827a25d0d7f8efc022eb4241802d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56a17336273171ffb8d00db9491656ee8aa618c318592ccc51096bb63e186adc +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..9cc4a5495602a609963ee3cbdb783ccab87151ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4ca1e15d89cd53f90c144c25a90218e8c2ecf0a1b5439e900bbe601aebf03dfc +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..e2df96257b7b29037b49b185540717f7640c96dd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0a81b0562d6e4cb7be8bfac7f85dd0325c86204e784caa1133023ff747843b17 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..fe6cfa68115707e65f08af42b81cdb5c8859b93f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76aa99a8b40b94035eee63425652f5d68399a9bd507deaf951fd4175d0fe7c77 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..f0f0d5b5375275381a9cac2741389b2073259d28 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/4_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.0240753889083862, + "learning_rate": 2e-05, + "loss": 0.209, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.773343563079834, + "learning_rate": 2e-05, + "loss": 0.4006, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.9890713095664978, + "learning_rate": 2e-05, + "loss": 0.3514, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.7700861692428589, + "learning_rate": 2e-05, + "loss": 0.4042, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.195196509361267, + "learning_rate": 2e-05, + "loss": 0.2426, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.6621294021606445, + "learning_rate": 2e-05, + "loss": 0.1404, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.247157573699951, + "learning_rate": 2e-05, + "loss": 0.3303, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.2074408531188965, + "learning_rate": 2e-05, + "loss": 0.5588, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.590269684791565, + "learning_rate": 2e-05, + "loss": 0.1526, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.6297767162323, + "learning_rate": 2e-05, + "loss": 0.2188, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.3473517894744873, + "learning_rate": 2e-05, + "loss": 1.0015, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.2105010747909546, + "learning_rate": 2e-05, + "loss": 0.4824, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.4897263050079346, + "learning_rate": 2e-05, + "loss": 0.3245, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.9941797256469727, + "learning_rate": 2e-05, + "loss": 0.4608, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.5566309690475464, + "learning_rate": 2e-05, + "loss": 0.3736, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.77103590965271, + "learning_rate": 2e-05, + "loss": 0.3665, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8090458512306213, + "learning_rate": 2e-05, + "loss": 0.212, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.363377332687378, + "learning_rate": 2e-05, + "loss": 0.5171, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.67446768283844, + "learning_rate": 2e-05, + "loss": 0.2245, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.057765483856201, + "learning_rate": 2e-05, + "loss": 0.6199, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.4825180768966675, + "learning_rate": 2e-05, + "loss": 0.3949, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.7534990310668945, + "learning_rate": 2e-05, + "loss": 0.2844, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.5341320037841797, + "learning_rate": 2e-05, + "loss": 0.3872, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.783373475074768, + "learning_rate": 2e-05, + "loss": 0.3564, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.5615463256835938, + "learning_rate": 2e-05, + "loss": 0.2161, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.308306932449341, + "learning_rate": 2e-05, + "loss": 0.7563, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.9441331624984741, + "learning_rate": 2e-05, + "loss": 0.2903, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.9952667355537415, + "learning_rate": 2e-05, + "loss": 0.1224, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.997939109802246, + "learning_rate": 2e-05, + "loss": 0.3494, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.912369728088379, + "learning_rate": 2e-05, + "loss": 0.405, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.1861718893051147, + "learning_rate": 2e-05, + "loss": 0.231, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.300227642059326, + "learning_rate": 2e-05, + "loss": 0.4797, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.1022934913635254, + "learning_rate": 2e-05, + "loss": 0.3365, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.8765240907669067, + "learning_rate": 2e-05, + "loss": 0.3469, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 4.114320755004883, + "learning_rate": 2e-05, + "loss": 0.3945, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.6984522342681885, + "learning_rate": 2e-05, + "loss": 0.2491, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.8632751107215881, + "learning_rate": 2e-05, + "loss": 0.2783, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.4998695850372314, + "learning_rate": 2e-05, + "loss": 0.5198, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.9596168994903564, + "learning_rate": 2e-05, + "loss": 0.3597, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.199298143386841, + "learning_rate": 2e-05, + "loss": 0.2839, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.5701171159744263, + "learning_rate": 2e-05, + "loss": 0.2495, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.7656409740448, + "learning_rate": 2e-05, + "loss": 0.3466, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.264856219291687, + "learning_rate": 2e-05, + "loss": 0.1516, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0180593729019165, + "learning_rate": 2e-05, + "loss": 0.2169, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.2818539142608643, + "learning_rate": 2e-05, + "loss": 0.3803, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.718665838241577, + "learning_rate": 2e-05, + "loss": 0.2722, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.792602062225342, + "learning_rate": 2e-05, + "loss": 0.4226, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.6356031894683838, + "learning_rate": 2e-05, + "loss": 0.3522, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.3408199548721313, + "learning_rate": 2e-05, + "loss": 0.0966, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.472543716430664, + "learning_rate": 2e-05, + "loss": 0.5309, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 6047770452426752.0, + "train_loss": 0.3530580139160156, + "train_runtime": 123.8719, + "train_samples_per_second": 3.229, + "train_steps_per_second": 0.807 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 6047770452426752.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..de65ae475ec00a6f9a02a85cdaad13a56cc5cea6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a230de2de749a98752bc7830e7815f58baf6098056892ec94232b45dc508334d +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..02f505af376216517dbf89b9d217c35b3f9f03d5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:05b5068ea9c75c3543bb9d5cdb042028e5a99ac55781586835559536f9a9a03a +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..3c30ee92416c785f777e7500636afec814f03523 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:58d0a1c74abede5d036d57a391aa8e98a8ef22bf0ea8b7421e7ddf048762dc0e +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..ecb96e0b037a9dbd82659df9fcfe5763b588380f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f872289734f93d126fd1aeac29ed56c4d071f0f0c5e1ae8df7be3453a4904e29 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..0873cef313c97e7ac10baa0363368e649db52fcf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cb86b0acda52c63b6d27c0c5f24487b2b6cae85369aba21dc95b5efb085ef16b +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d0fabf752538d41e0859e8c1a6ab6e5088d24a8a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2405d0b23a46a2eb8e6c5f0a88477c2e707bf2b8c0c5228dc22d766976b24232 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..484ad6e2d1a6d1dcac4c77d5a108460ff4a380cf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:32bd90a19104f5a8e50eb2d54daa2ca0ac14b6b85b3617fba51157ee967dd5ef +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..a204abd7dd19082ac46b75b72b99f6324453b2fc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5bd49e7edf03fb0f8e08129cb1df532731143fe857835c501c39597662ce948e +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..64613e55208fae40775fc99835c98f0dca330ab4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/5_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.4192941188812256, + "learning_rate": 2e-05, + "loss": 0.1753, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.7077388763427734, + "learning_rate": 2e-05, + "loss": 0.3513, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.4440271854400635, + "learning_rate": 2e-05, + "loss": 0.1267, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.11960666626691818, + "learning_rate": 2e-05, + "loss": 0.1086, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8296095132827759, + "learning_rate": 2e-05, + "loss": 0.1107, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.26236826181411743, + "learning_rate": 2e-05, + "loss": 0.0536, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.015287407673895359, + "learning_rate": 2e-05, + "loss": 0.0326, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.13008597493171692, + "learning_rate": 2e-05, + "loss": 0.0277, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.418893337249756, + "learning_rate": 2e-05, + "loss": 0.0827, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.7668967247009277, + "learning_rate": 2e-05, + "loss": 0.1154, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.0926952362060547, + "learning_rate": 2e-05, + "loss": 0.0917, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.5399396419525146, + "learning_rate": 2e-05, + "loss": 0.136, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.6366981863975525, + "learning_rate": 2e-05, + "loss": 0.2528, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.1575629711151123, + "learning_rate": 2e-05, + "loss": 0.631, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.13304385542869568, + "learning_rate": 2e-05, + "loss": 0.0369, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.9152341485023499, + "learning_rate": 2e-05, + "loss": 0.0433, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.9383106231689453, + "learning_rate": 2e-05, + "loss": 0.0849, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.30875903367996216, + "learning_rate": 2e-05, + "loss": 0.0428, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.2635672092437744, + "learning_rate": 2e-05, + "loss": 0.0483, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.0975215435028076, + "learning_rate": 2e-05, + "loss": 0.1204, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.2128736525774002, + "learning_rate": 2e-05, + "loss": 0.1653, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.4516749382019043, + "learning_rate": 2e-05, + "loss": 0.2333, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.2670416235923767, + "learning_rate": 2e-05, + "loss": 0.224, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.22837549448013306, + "learning_rate": 2e-05, + "loss": 1.0722, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.764431476593018, + "learning_rate": 2e-05, + "loss": 0.2289, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.295931339263916, + "learning_rate": 2e-05, + "loss": 0.1402, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.22568394243717194, + "learning_rate": 2e-05, + "loss": 0.0229, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.517451047897339, + "learning_rate": 2e-05, + "loss": 0.1174, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.2220073789358139, + "learning_rate": 2e-05, + "loss": 0.0146, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5590463876724243, + "learning_rate": 2e-05, + "loss": 0.0357, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.41243040561676025, + "learning_rate": 2e-05, + "loss": 0.0399, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.112615585327148, + "learning_rate": 2e-05, + "loss": 0.5905, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.7114840149879456, + "learning_rate": 2e-05, + "loss": 0.0613, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.27175503969192505, + "learning_rate": 2e-05, + "loss": 0.0327, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.5079922676086426, + "learning_rate": 2e-05, + "loss": 0.2761, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.504249095916748, + "learning_rate": 2e-05, + "loss": 0.3168, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.11905453354120255, + "learning_rate": 2e-05, + "loss": 0.2349, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.19908714294433594, + "learning_rate": 2e-05, + "loss": 0.0109, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.09087877720594406, + "learning_rate": 2e-05, + "loss": 0.5035, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.1647261381149292, + "learning_rate": 2e-05, + "loss": 0.0599, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.48274871706962585, + "learning_rate": 2e-05, + "loss": 0.0642, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.13012978434562683, + "learning_rate": 2e-05, + "loss": 0.0752, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.27074968814849854, + "learning_rate": 2e-05, + "loss": 0.0393, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.114485502243042, + "learning_rate": 2e-05, + "loss": 0.0555, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.4042900800704956, + "learning_rate": 2e-05, + "loss": 0.0416, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.757662057876587, + "learning_rate": 2e-05, + "loss": 0.2188, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.598297119140625, + "learning_rate": 2e-05, + "loss": 0.0291, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.773273229598999, + "learning_rate": 2e-05, + "loss": 0.1819, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.053060125559568405, + "learning_rate": 2e-05, + "loss": 0.2543, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 5.627431869506836, + "learning_rate": 2e-05, + "loss": 1.329, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289199494234112.0, + "train_loss": 0.1868517017364502, + "train_runtime": 119.1125, + "train_samples_per_second": 3.358, + "train_steps_per_second": 0.84 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289199494234112.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..4c82b234c0623ecedaabb11959faace6769298cf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:494ba77f631cb339181aef2e0912002b77c416c51aeac505d83dfaf42134bddc +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..58838f4cef64a74eef17e14128b1f2a00556d6b1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:967d8dc8bb946445d2492a814c903c73e18985654cded6a753f7372f209ad9e5 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..c748e1fc1915a63ed68fb9c64d2da14e32c58add --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d9f86e27eb260ff29d1bd08fc96b6b9a5e19d5fbd4625af6ef05c6be88b4537 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..551998c491ef51584fa466cb0420ca3eb400c6f6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4c64ba131b3197b51464a952228ca9fb1242756f131b0fd6bc031c8a16c96ab6 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..c2e60212ba2c0453e8f4c027c1213088048617da --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:240d8f4473cf5251c6d0cad9b137936e349180d77c35a49a2b1f77b0029dd769 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..dc4bb4ffea773c2724c0383563e3999ec0da867f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d7567471f940f975aa8f7832dbbf18e562f51d9ceb33388928cc3a057e7af9ff +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..61b80acd329a29f01ec855b73184785ae90ced47 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:358421098b4cd0765e25759c3d12e3478417c64b0a29190fce2dd51a3aeaa9b1 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f5642a29e93286beb9f54410eb502af582db922f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fd9beddfaa5389bced1db17bbd05daa839d360a63a44d749b6d3fb35bd3a4e65 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..77831627cc970432bb185681cef908bcc821f510 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/6_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.736836910247803, + "learning_rate": 2e-05, + "loss": 0.4038, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.707756042480469, + "learning_rate": 2e-05, + "loss": 0.42, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.8180034160614014, + "learning_rate": 2e-05, + "loss": 0.3745, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.271767258644104, + "learning_rate": 2e-05, + "loss": 0.5181, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.852658748626709, + "learning_rate": 2e-05, + "loss": 0.4805, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.47449979186058044, + "learning_rate": 2e-05, + "loss": 0.0357, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.786670446395874, + "learning_rate": 2e-05, + "loss": 0.3325, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.331701278686523, + "learning_rate": 2e-05, + "loss": 0.507, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.011174440383911, + "learning_rate": 2e-05, + "loss": 0.7319, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.263643503189087, + "learning_rate": 2e-05, + "loss": 0.3701, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.9021544456481934, + "learning_rate": 2e-05, + "loss": 0.3514, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.5527875423431396, + "learning_rate": 2e-05, + "loss": 0.3056, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.1424427032470703, + "learning_rate": 2e-05, + "loss": 0.3729, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.804325580596924, + "learning_rate": 2e-05, + "loss": 0.2614, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 4.539284706115723, + "learning_rate": 2e-05, + "loss": 0.8229, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.942174196243286, + "learning_rate": 2e-05, + "loss": 0.3593, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.0257210731506348, + "learning_rate": 2e-05, + "loss": 0.3943, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.8576889038085938, + "learning_rate": 2e-05, + "loss": 0.5704, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.6807342767715454, + "learning_rate": 2e-05, + "loss": 0.3485, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.2396899461746216, + "learning_rate": 2e-05, + "loss": 0.6077, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.9702136516571045, + "learning_rate": 2e-05, + "loss": 0.8195, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.9754537343978882, + "learning_rate": 2e-05, + "loss": 0.2183, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.770653486251831, + "learning_rate": 2e-05, + "loss": 0.4191, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.682520866394043, + "learning_rate": 2e-05, + "loss": 0.3175, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.384674310684204, + "learning_rate": 2e-05, + "loss": 0.4137, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.3491796255111694, + "learning_rate": 2e-05, + "loss": 0.3945, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.5511598587036133, + "learning_rate": 2e-05, + "loss": 0.3926, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 3.1397528648376465, + "learning_rate": 2e-05, + "loss": 0.8281, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.5708088874816895, + "learning_rate": 2e-05, + "loss": 0.4833, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.441242218017578, + "learning_rate": 2e-05, + "loss": 0.5706, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.5901715755462646, + "learning_rate": 2e-05, + "loss": 0.3286, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.8951689004898071, + "learning_rate": 2e-05, + "loss": 0.4622, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.1505897045135498, + "learning_rate": 2e-05, + "loss": 0.5493, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.1398003101348877, + "learning_rate": 2e-05, + "loss": 0.3834, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.0591495037078857, + "learning_rate": 2e-05, + "loss": 0.6738, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.0965842008590698, + "learning_rate": 2e-05, + "loss": 0.3464, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 6.196214199066162, + "learning_rate": 2e-05, + "loss": 0.7417, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.7029244899749756, + "learning_rate": 2e-05, + "loss": 0.5587, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.9107545614242554, + "learning_rate": 2e-05, + "loss": 0.2592, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.707282543182373, + "learning_rate": 2e-05, + "loss": 0.1958, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.8375647068023682, + "learning_rate": 2e-05, + "loss": 0.4595, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.7143489122390747, + "learning_rate": 2e-05, + "loss": 0.5923, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.834902286529541, + "learning_rate": 2e-05, + "loss": 0.2336, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.2900763750076294, + "learning_rate": 2e-05, + "loss": 0.4355, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.1141912937164307, + "learning_rate": 2e-05, + "loss": 0.291, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.8684674501419067, + "learning_rate": 2e-05, + "loss": 0.9307, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.7162550687789917, + "learning_rate": 2e-05, + "loss": 0.3088, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.837238073348999, + "learning_rate": 2e-05, + "loss": 0.3638, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.3492848873138428, + "learning_rate": 2e-05, + "loss": 0.4648, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.920907974243164, + "learning_rate": 2e-05, + "loss": 1.0586, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 1.0467000203608064e+16, + "train_loss": 0.4612684631347656, + "train_runtime": 179.0074, + "train_samples_per_second": 2.235, + "train_steps_per_second": 0.559 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 1.0467000203608064e+16, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..b4df29a0f51fdfcc4a339837a84a2678e9142079 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:62b13cb0a7f6f85af8ffa406b685db53c9fdf38bd1749c78fe0f68fba63bd422 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..7a80cf828cde0a0c5437cb0f3b46258872a899cd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:01d05955f126bb76f6b808f67cbb45b36ac563fff39206378e9eab4fd1e623e9 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..0e27419108b05ab179563cbf95fd9b7eb210d57c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cc8bd18d64ea4a51266e0be35cfaff7d078dc87e32cce9998890cf56857cc6c8 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..aa711529017e34a53b64680646816d1e8fa929d2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c0ac954468f229a8d3108e6a9755a4519c6207dc919dcc695c63dea7caa96ac +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..e9cae803d7234922c692d2dbac1a9d270fb054ee --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:48346704c4732edcc73928a0c99da7abd1b4bfb50b6c2427fa2eeea4633427c6 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..73408f00b4ac4170999818ac59de3418e98d7ddd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d6cb3ef05e58e25915782597cf783c990a91181fb67a27660b06a0a4d9e207f +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..3c8eb39f602503e83392344240ddc15dfea1fc27 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:26466697b35443abb1e74bba7a42697fff2d8421ea9a3adfcf2321c2363fdd57 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8147dbb5620b4f36d703cb04caa98ef10282e032 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c321361cf222ab7b76cfdf64e9009e44de07187d1a90a47eca109101e605bcd9 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..dc0173ae61fb9150e06f0949200517d81dc45cd3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/7_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.3207453489303589, + "learning_rate": 2e-05, + "loss": 0.0712, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.7498466968536377, + "learning_rate": 2e-05, + "loss": 0.3689, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.7817299365997314, + "learning_rate": 2e-05, + "loss": 0.1917, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.31601881980896, + "learning_rate": 2e-05, + "loss": 0.2284, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.532356023788452, + "learning_rate": 2e-05, + "loss": 0.6504, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.13659778237342834, + "learning_rate": 2e-05, + "loss": 0.0347, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.6198537349700928, + "learning_rate": 2e-05, + "loss": 0.0538, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.61257004737854, + "learning_rate": 2e-05, + "loss": 0.4736, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.297581195831299, + "learning_rate": 2e-05, + "loss": 0.2017, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.916984796524048, + "learning_rate": 2e-05, + "loss": 0.2665, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.14377397298812866, + "learning_rate": 2e-05, + "loss": 0.194, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.438236951828003, + "learning_rate": 2e-05, + "loss": 0.2723, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.48461809754371643, + "learning_rate": 2e-05, + "loss": 0.0278, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.6433966159820557, + "learning_rate": 2e-05, + "loss": 0.2991, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.43798476457595825, + "learning_rate": 2e-05, + "loss": 0.298, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.7849429249763489, + "learning_rate": 2e-05, + "loss": 0.3026, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.4565381109714508, + "learning_rate": 2e-05, + "loss": 0.0272, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.4069461822509766, + "learning_rate": 2e-05, + "loss": 0.4543, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.8650145530700684, + "learning_rate": 2e-05, + "loss": 0.0605, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.771399021148682, + "learning_rate": 2e-05, + "loss": 0.673, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.4950507581233978, + "learning_rate": 2e-05, + "loss": 0.4736, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.520219087600708, + "learning_rate": 2e-05, + "loss": 0.2502, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.930126667022705, + "learning_rate": 2e-05, + "loss": 0.3065, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.17428846657276154, + "learning_rate": 2e-05, + "loss": 0.0079, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.5029134154319763, + "learning_rate": 2e-05, + "loss": 0.4142, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.0570330619812012, + "learning_rate": 2e-05, + "loss": 0.4269, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.03667684271931648, + "learning_rate": 2e-05, + "loss": 0.3434, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4633074998855591, + "learning_rate": 2e-05, + "loss": 0.0261, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.12067744135856628, + "learning_rate": 2e-05, + "loss": 0.1116, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.948736667633057, + "learning_rate": 2e-05, + "loss": 1.2452, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.9964993596076965, + "learning_rate": 2e-05, + "loss": 0.0726, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.145695924758911, + "learning_rate": 2e-05, + "loss": 0.5176, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.6078904867172241, + "learning_rate": 2e-05, + "loss": 0.1524, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.5305264592170715, + "learning_rate": 2e-05, + "loss": 0.127, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.972503662109375, + "learning_rate": 2e-05, + "loss": 0.3565, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.3434875011444092, + "learning_rate": 2e-05, + "loss": 0.2697, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.105001926422119, + "learning_rate": 2e-05, + "loss": 0.4565, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.45210543274879456, + "learning_rate": 2e-05, + "loss": 0.3477, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.812847375869751, + "learning_rate": 2e-05, + "loss": 0.3956, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.900662899017334, + "learning_rate": 2e-05, + "loss": 0.8243, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.3302174806594849, + "learning_rate": 2e-05, + "loss": 0.1611, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.6209325790405273, + "learning_rate": 2e-05, + "loss": 0.1114, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.9477224349975586, + "learning_rate": 2e-05, + "loss": 0.4568, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.0254838466644287, + "learning_rate": 2e-05, + "loss": 0.3511, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.6387124061584473, + "learning_rate": 2e-05, + "loss": 0.3199, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.8599830269813538, + "learning_rate": 2e-05, + "loss": 0.1333, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.1979265213012695, + "learning_rate": 2e-05, + "loss": 0.1062, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.2699570953845978, + "learning_rate": 2e-05, + "loss": 0.0708, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.4119693040847778, + "learning_rate": 2e-05, + "loss": 0.1798, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.1397922039031982, + "learning_rate": 2e-05, + "loss": 0.2951, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5500391349288960.0, + "train_loss": 0.2892104434967041, + "train_runtime": 123.2936, + "train_samples_per_second": 3.244, + "train_steps_per_second": 0.811 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5500391349288960.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..7a92ac2ddae291d825312d94a692c88e41d3fd4e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d8d1a9251da24a2463cd70ac2925690985c18606580b64020b88688be728db81 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..5696e80c708f5ddc0499f2a66d25f72f4fbe434d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb6df497d50abca59aa4f5016dda34f353ee99ecdb3a85564fffa603312703c2 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9e6ec0c590c4d95b733d6e33e37d733e6c4e558 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8eaee0cb3dacb72b8068bbb5d3383451ee102e02f8348f613005a57b2f90a989 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..640ee95965a57042eb05c0cc76c9df695e484c9a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b86b3c9d3ee19d96687b6fb1fafccd74258d6bfa448a6c197ede424348927e3 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..3866448d9a546128a3f9ed0605562a3e6402a86a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:928f4db3262a6f933f27f7b957e47669f74a3f37448de7b633f66e511dd90d86 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..46508710c0264dffac700ba4e312f989a480728c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8f2a14f6e78f58937c9176ed49aafeee311f09a8842d8c7ae7a813b058b16d47 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..fa025ad58026db9363600cb0c5c80fb8a3f3fcc9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f0b9e1692e2d7358d8bca5e2d3aea98928c6ab7fa68aec4e429656b429f9a1c6 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..be548d96d4c2d4b36630b05acd97142e8891973e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c30985386c7bdfc653141e8fa7ca5ca3785c7a1d1b0752b75f3e35f4c3eb0f75 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..ff7368762d23ba9cae0531bbc915cda69425189b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/8_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.05183964595198631, + "learning_rate": 2e-05, + "loss": 0.0248, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.08121512830257416, + "learning_rate": 2e-05, + "loss": 0.0941, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.448318749666214, + "learning_rate": 2e-05, + "loss": 0.0482, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.4191150665283203, + "learning_rate": 2e-05, + "loss": 0.2959, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.04907870292663574, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.3470380306243896, + "learning_rate": 2e-05, + "loss": 0.0718, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.07462665438652039, + "learning_rate": 2e-05, + "loss": 0.1667, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.852299690246582, + "learning_rate": 2e-05, + "loss": 0.3256, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.5574058294296265, + "learning_rate": 2e-05, + "loss": 0.1019, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 4.223021030426025, + "learning_rate": 2e-05, + "loss": 2.2283, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.453669548034668, + "learning_rate": 2e-05, + "loss": 0.4525, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.042588233947754, + "learning_rate": 2e-05, + "loss": 0.3349, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.3411954641342163, + "learning_rate": 2e-05, + "loss": 0.0612, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.9152738451957703, + "learning_rate": 2e-05, + "loss": 0.0458, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.1214474439620972, + "learning_rate": 2e-05, + "loss": 0.0411, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.046567466109991074, + "learning_rate": 2e-05, + "loss": 0.0078, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.626150369644165, + "learning_rate": 2e-05, + "loss": 0.7565, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.272691488265991, + "learning_rate": 2e-05, + "loss": 0.3312, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.244640827178955, + "learning_rate": 2e-05, + "loss": 0.6516, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.49138605594635, + "learning_rate": 2e-05, + "loss": 0.6353, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 6.054605484008789, + "learning_rate": 2e-05, + "loss": 0.7234, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.0022809277288615704, + "learning_rate": 2e-05, + "loss": 0.0752, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.6537396907806396, + "learning_rate": 2e-05, + "loss": 0.4751, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.8570873737335205, + "learning_rate": 2e-05, + "loss": 0.2354, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.2755966782569885, + "learning_rate": 2e-05, + "loss": 0.2215, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.04173465073108673, + "learning_rate": 2e-05, + "loss": 0.1996, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.4190864562988281, + "learning_rate": 2e-05, + "loss": 0.2807, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4271361827850342, + "learning_rate": 2e-05, + "loss": 0.2469, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.3003133535385132, + "learning_rate": 2e-05, + "loss": 0.0695, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09996131807565689, + "learning_rate": 2e-05, + "loss": 0.0089, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.4211537837982178, + "learning_rate": 2e-05, + "loss": 0.1875, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.459202766418457, + "learning_rate": 2e-05, + "loss": 0.6795, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.1551332473754883, + "learning_rate": 2e-05, + "loss": 0.095, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.8788769245147705, + "learning_rate": 2e-05, + "loss": 0.031, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.43607568740844727, + "learning_rate": 2e-05, + "loss": 0.0201, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.7366693019866943, + "learning_rate": 2e-05, + "loss": 0.193, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.0499285459518433, + "learning_rate": 2e-05, + "loss": 0.0474, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.562388002872467, + "learning_rate": 2e-05, + "loss": 0.0383, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.1588374823331833, + "learning_rate": 2e-05, + "loss": 0.1371, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.2984813451766968, + "learning_rate": 2e-05, + "loss": 0.0254, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 4.474546909332275, + "learning_rate": 2e-05, + "loss": 0.5605, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.43176889419555664, + "learning_rate": 2e-05, + "loss": 0.0199, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.23999084532260895, + "learning_rate": 2e-05, + "loss": 0.0124, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.7237604856491089, + "learning_rate": 2e-05, + "loss": 0.0386, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3866379261016846, + "learning_rate": 2e-05, + "loss": 0.1693, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 5.821412086486816, + "learning_rate": 2e-05, + "loss": 0.5902, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.24462923407554626, + "learning_rate": 2e-05, + "loss": 0.0095, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.447963237762451, + "learning_rate": 2e-05, + "loss": 0.1376, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.23387965559959412, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.16078920662403107, + "learning_rate": 2e-05, + "loss": 0.0097, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5327278363901952.0, + "train_loss": 0.24452131986618042, + "train_runtime": 116.6146, + "train_samples_per_second": 3.43, + "train_steps_per_second": 0.858 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5327278363901952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..9cd06283e4b26552f6d7bc5f6198837568582c4d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4938ecaf14e664e500f2d6fad072f53a9cb1d86ce487da48762ff04ab1503e35 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..977e45fba1aa62ca45f8262167b27ab08e2da9ff --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d036db5417d2878013f2d8942c4fbb4cb6f9b0899079e4235664639eb2d26a68 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..016eea15d6cc3ede890ae9130e6af05a8d2d9553 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0110c47218982f6436af680553d4ea6061303f7a25f8836484607818e3ddc377 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..3a1122584e4a80a2af1535162ebd8894a2acda31 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eac92657bc0038f07499757c3acde96b6fb75fe29b4d8ebd85ff0c3d32637b10 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..3aa342dbc43fae8900d6a46a34c72d6f242447f5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3599c7e4144cdad2a40ae7a2bc83445ea2be23a11401d0dab4f875c99941fd95 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..89776f9efc717731f710edc56760cd9b97c33f58 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:90af3efe6f38ebabcc7e58822031c1972e475c22cb457689490a28aa0bfe768c +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..bf03df88172b49f95ca9ea5d775b1d142b6296e5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:092cc05b489972ed25c20a5ce18d9029d6915ad5368d4b55808ac41c01afa796 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..a392c1ec934d6d9bf67d28b4556a5e206407c33e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5bc9f172fe896eee0024703f49ac16998f2da716529600c9e61bb3720b176df9 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3b110e0da28faef8f2ab79e2904f67be4af0da1a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/9_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.28986603021621704, + "learning_rate": 2e-05, + "loss": 0.3234, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.6648087501525879, + "learning_rate": 2e-05, + "loss": 0.7266, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.6703162789344788, + "learning_rate": 2e-05, + "loss": 0.0662, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.9952301979064941, + "learning_rate": 2e-05, + "loss": 0.1477, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.16450397670269012, + "learning_rate": 2e-05, + "loss": 0.351, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.7219018340110779, + "learning_rate": 2e-05, + "loss": 0.1621, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.060893815010786057, + "learning_rate": 2e-05, + "loss": 0.3696, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.3581697940826416, + "learning_rate": 2e-05, + "loss": 0.6065, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.739429235458374, + "learning_rate": 2e-05, + "loss": 0.1457, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.182863235473633, + "learning_rate": 2e-05, + "loss": 0.2312, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.7728469371795654, + "learning_rate": 2e-05, + "loss": 0.7462, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 4.030282020568848, + "learning_rate": 2e-05, + "loss": 0.4916, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.216357707977295, + "learning_rate": 2e-05, + "loss": 0.4823, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.2622979879379272, + "learning_rate": 2e-05, + "loss": 0.0623, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.1975100040435791, + "learning_rate": 2e-05, + "loss": 0.2119, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.11730964481830597, + "learning_rate": 2e-05, + "loss": 0.0047, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.9041928648948669, + "learning_rate": 2e-05, + "loss": 0.5413, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.8735140562057495, + "learning_rate": 2e-05, + "loss": 0.165, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.37345966696739197, + "learning_rate": 2e-05, + "loss": 0.0605, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.45941162109375, + "learning_rate": 2e-05, + "loss": 0.3434, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.4740889072418213, + "learning_rate": 2e-05, + "loss": 0.2968, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 5.144214153289795, + "learning_rate": 2e-05, + "loss": 0.6291, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.8222898244857788, + "learning_rate": 2e-05, + "loss": 0.0589, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.9600136280059814, + "learning_rate": 2e-05, + "loss": 0.381, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.699195146560669, + "learning_rate": 2e-05, + "loss": 0.0662, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.872018039226532, + "learning_rate": 2e-05, + "loss": 0.09, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.7105398178100586, + "learning_rate": 2e-05, + "loss": 0.4963, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.1588001698255539, + "learning_rate": 2e-05, + "loss": 0.0101, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.2687864303588867, + "learning_rate": 2e-05, + "loss": 0.3525, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5537347793579102, + "learning_rate": 2e-05, + "loss": 0.097, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 3.1914584636688232, + "learning_rate": 2e-05, + "loss": 0.2065, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.6231799125671387, + "learning_rate": 2e-05, + "loss": 0.2172, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 4.902224540710449, + "learning_rate": 2e-05, + "loss": 0.2503, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.0578718185424805, + "learning_rate": 2e-05, + "loss": 0.0819, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.19334405660629272, + "learning_rate": 2e-05, + "loss": 0.0184, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 6.02439546585083, + "learning_rate": 2e-05, + "loss": 0.552, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.0188320875167847, + "learning_rate": 2e-05, + "loss": 0.2332, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.2132835388183594, + "learning_rate": 2e-05, + "loss": 0.3154, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.4453291594982147, + "learning_rate": 2e-05, + "loss": 0.1142, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.5968852043151855, + "learning_rate": 2e-05, + "loss": 0.3602, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.364363670349121, + "learning_rate": 2e-05, + "loss": 0.2886, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.0657949447631836, + "learning_rate": 2e-05, + "loss": 0.1968, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 3.505636692047119, + "learning_rate": 2e-05, + "loss": 0.2307, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.9454941749572754, + "learning_rate": 2e-05, + "loss": 0.3864, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.423117995262146, + "learning_rate": 2e-05, + "loss": 0.2769, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.70637845993042, + "learning_rate": 2e-05, + "loss": 0.2918, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.083464622497559, + "learning_rate": 2e-05, + "loss": 0.5557, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.45639291405677795, + "learning_rate": 2e-05, + "loss": 0.0504, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.2392776012420654, + "learning_rate": 2e-05, + "loss": 0.2418, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.2385011613368988, + "learning_rate": 2e-05, + "loss": 0.057, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291656215527424.0, + "train_loss": 0.27284836292266845, + "train_runtime": 117.487, + "train_samples_per_second": 3.405, + "train_steps_per_second": 0.851 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291656215527424.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..deebaa9367868cb88ef2899c8f192b63ce3c25d1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b2a96e971027bcee5035578c5c103d43e3652ebacb6fff40645651b9582d042 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d836df2492e460917486ed8136e0b0f3e3fecf5b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc63bd13932ed35e76a80203bbb683008c34cd20bf6a6892d3a0a84edb5d1a72 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_0/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..adfdd070d26a2bb9946837ff7fa6afbb2d73b8e1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:540aeb07faa350d65024d9e91a4ce58ab2ad1a7d28e7551bcdd6f777b21d8192 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..acd34681374adb5f4e00be9d954bf5ab9a55af44 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e363c3ab92f4afd84d075e1503c75f16fe186f602ee7bf8ff8191eaf9bdad0c7 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_1/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9179833308311147eebe43c696fbfa735dc4a480 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b115ebdb018f700e89ac6ed653f3754cad5e961a9c0c0b96e5239becc2333265 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9b2c23f54e62fe442df7aea7b2a759333e338dd2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cbbfe989725e2832e6446c6ed8b1eaf90ea716bcf6057d2aa07ea342aee7bed1 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_10/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9f652b571c535e5e1887d1c1cb3faf816bf7a858 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c4de79116cdd449ed49926a923f70b5cffe434d9f2375fa7831372149c184d9 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..dfd572ae56bb9d7329ee6cc776889d7656b28442 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ea6f3e854cc6636f6be3447ec01b1d4f9c728247cfc2739cdc740875050635db +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_11/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..54063740f5e9e268ae51fec11f594065ef5dcac1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09985a79c9ec72348ddd48d116a07b384efce949e5f018ff03b19d253c843885 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c217c4f6b8303d9230f78dbf5829d51ce00c5a38 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7f5628250998ed263a9c6a0c693d7d3c4a014575196efa6b0169796e1f606ab7 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_12/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6f61d75f364cfcc9ab87dcc54a922814bdeb1be0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9071c9349acbda14efa217369580d5867e20dcd0290db692dd224af197c9ebe1 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a5caf4feb217752ff32faa687d48c773b09a5876 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:48f1345eb4c51dd543d1f08a1b6b4da1a870a4cd7e7df9ffaf44588f7e7000a1 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_13/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..acad1776bb501d870c2bf2ecc3cec71406762572 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1a24d39720d8bc7e97cc60a2ae641c81debfcec2cdafbd2ba3a7a69833e1279c +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..78a672b179c21121dc1c199f9cc548edc3ee5725 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:38081e75ab68435cd7ce4589a93982063e6314b8425a515df60ec56205275824 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_14/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..07f3155b4de1f51b6651a2dbe5075ac0f0ffab0f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ef6da2718d3bf9dd624d40900efe44bf8d31d078cab27795d47e3be7232374b8 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1e3317c58bd5d9596c58ede464defd02a6a7c9a4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f439120c7b7b867de6bed55ed08bc1f7ce79b65d866b71cee8248d3bf6702df6 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_15/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f80ffab8e6d95bf79482b83b1cdf3f92c769656d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:14965aef6168e1848b1f12a9dbcdaf8c84b0efe9fdd383adc6e77537f5bd9ace +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ff5467127d1855eb3d98828f3b1c10aad3f8895a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d3858128482e9d701677c440dac0b6e01ac052b4619e304e4af206f1fc356b9c +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_16/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e6ef8fbdbcd5f763be6112f21993df66a27d28a3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:229df355df3dfdcbcf6d04fdace08ac691c0dfebe70a4305560a0141a527c554 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f66a5b77fdeb0f7a7f99ccd06c4c5c7ed4b22c16 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8b0ca4cc8bbd106d52801efdd6724dfa6395198e2601fbbd214f17c5375d80e2 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_17/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..cd9f2ca33e813857328d54aa487ac97e0d16ebfb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e33ffd75ad2a8bf07f21a4430693019e0a06da37b2821cb60a52e20b261efe56 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7079c01ecaca8c231e54043cdeea2b68c1b872ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6d9a1446969325c50b87ae2970274d0b074e2048df186a85fda34d039dfa1935 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_18/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..68e1f99bc460b647fc6b50e86ee201718c9bcc90 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7bde6cd5e1b5eb0d92e49797ed4762adb3cadf734aa4c1362e4c64a170252e9a +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..8be64cc38cfc41ed8bfc5d30758b9a4084264b52 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6c0aada2385aba2548047841bceb0495906317ad5f25b2b60b8979e65c2dbf1a +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_19/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..bddb2e587d254f1b1cb0311e4deeec604eb8468a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b0c5c85b5a805e946791bedc4b3fce4e7277550b28b6da08391d02287e85c391 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a29bc6b7177bfc78dfb42499d963c13e3a5ad806 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9b42d313d82dad3d652b6cea6d2522bfe57ca92a8505ed46b36959a12d35d275 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_2/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..afb40a1a1ccb6a3adfb253cb1a6a392bfd106a00 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c0876d6f11765fa46c2fd35541cef08d7974fdd107fa6777189cf2ea0dd8d57a +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..25a7821ccf2369b4543fb76eb2218f89dff7832d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fdc10cb1e606421ec35c5c7b6563ffa80fd75e3000d4bddb5ad81eb752a53f6d +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_20/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..180d105bcf7999cf95a6dc4245580823274027b5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e7cd91074720e9a6ba078214040f781bf168310248ce233c73a715ca778d8038 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..5af1dcdbf7194e04ed108f7d1c7939f428a7b951 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:adbc1919f95de27647450ec76c067fed04c1a5329d46d4525c8470c09c07b700 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_3/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a2d08390020066f225034bf188918d09717882c4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7352f4be87efd2a55a5434fd37f5f4eaa889e52f5e4520eb73c0eb4a139b4810 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0a5b2a59718b5981a020b26f8a30a86290bbe7d9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:27160bd817c11369a0fafebd550b2c57a1df5574c6ca71f5e21975b9d4529c55 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_4/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..189d03f22e5c8d2de6ec88abba0098bff617a02d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:32674ec6af8ee06f00cb6a5048a9673324964d0623f3f3cec7ce74d238fd1c39 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..685eecdf4407b99183f4b6c433cbcd10ff25803d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76ea7159ce2c9ff6cce77fb15e05b2246d42cbce70093731a7494ab6933bac20 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_5/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..feff33636bff9734b0adaf3337185b3d2621ad73 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f59b72d479e84b35c8dd15aa90ca0f5e46c7998bfed8dc5865973892bc645836 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0cc8ea948759c92c9ece5b830d09583e0d199191 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:73b37b0851ee993860c32f5a25b9f5c33ee169c3af81dbd6c1e2f7ed5972979d +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_6/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6812d09dd5483d1e0d4f524c9263a4a65f5e910e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cbd81782d53e3c7c655aebd062b48fda7702c2665c70bb271eb0a9e58d1e780b +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6a708bfb3a34cf0df5c862ab29676eab5694f518 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e5218a5b3fe105e8c628dfd5081357418eec04d684e846d40f0b7bfa104040ed +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_7/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..551543731c937b5ad26409215c65d154c37145c1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:85a36b16f9256e8f7ff863de2d86e76eab0d401cf3db279809aa7983b5a6b591 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..502162aaf58f26269ca7b08a40dd0e9a45f44428 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:82fd71ca02a88374ae04b00a65aabca06ad96cf642c605e67cb608eadde018c5 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_8/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ba92af6582c920072347d8cc5b5d9c155b717ccc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cdd3a558d007bf2d97842036b31e6a3b68c0630a581162baaf49ced8fd0828c6 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..3713be5ce97dfc7534f46bd9be0fdff97c0e0775 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:69b68550ebafe421e6e29ad291c4d3e8e7a8008de696f6d13b1452c809328788 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/client_9/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round10_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round10_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..8480f3290fe02dcc8ad3f411dce4bb771607f136 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round10_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2a7a8d47edc80b205a0230dd0d312c16f3cc81aa0b01bd0d21718b6314735bb +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round11_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round11_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..120369764554a4d6cac58078d8512bf359607829 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round11_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf1dde35905b1ad125edf778aff1f9c3576ae0715c8c9876f665a63efcf18cd +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round12_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round12_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..64b454d945aa2a842c509c5cebe5ce0706d83a0c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round12_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb23ed1efc2080025ddd7ca376ff771b3f058afa24e048f5cb3cf99201aa1e03 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round13_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round13_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ad8d07954c41361e73a0c9c7d3b4b46d894b188 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round13_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8caca9dd439d57dbc5054f65d33ebbd7ac694940b2b9dbe5e8cdf68881137a2 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round14_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round14_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f05da1a0577050cea961f470b13842ab5f1d3f9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round14_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafca5706ac84f575d7648ee4786543561523e58cea8c01d2214b03d300bb988 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round15_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round15_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..351a4d2a284717aec90435f50067e1bb7002e9e4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round15_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:341a5a225c4c17a621d4f60caf4c34834000ea74e6f3c8b4f5aaada72ca1f547 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round16_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round16_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..08cfb19e58b68e297c605b87165d4acd241054e5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round16_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb378347913e2dd7be50fa6f9e3bd9e481f3897c2c9d31c597fe9164efc50bee +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round17_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round17_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..aeee08397c6123d33648f6703464379a4b7f85fe --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round17_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70c677ae195453774b8815440fee856ddb773c322681ac422ebc6e8abd5a58d7 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round18_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round18_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9935621fdc6da9e4abad379d97303b5ebf4f8f5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round18_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:393c3d5a15b40b2ff049c254f5b09f25d226b0971827b4e7460ca58f9ff7b76f +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round19_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round19_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..2766cf1a4da1390df01ea4146b4a3a7a06e1434e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round19_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dc77db94ac5b4746eebf2c850a9425da9a425a71108790d6b7f5f2654b614a8 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round1_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round1_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..34bfc5121e29fc587c3f574c17d8006189d3021a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round1_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47a94b88b812f21101345aae20a53cc10e0b346660c2cef669c77b201d06d508 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round20_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round20_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..4be93ba48f85d89c8cbd0e7426b79a10f51bc786 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round20_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20b9a2e67fdcb4051cebf52cba2ed0d47facb4dc69170f35a3483e3601bf25a4 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round2_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round2_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..b42d717305b30cd186c1c408b8719020c0658a5d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round2_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fd7458877767da0a36b80344cd6ec7eae8fbc8e11d07f8bfce3b4d5f5f239a9 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round3_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round3_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..01842d1316f21c3326c5735c6257b499da7ca458 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round3_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e537f056156082569ae0819abd06beb04b46581d0c63843bd68926e0ac2bff0e +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round4_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round4_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..3af9fe115499486166c3ef6d0a7bd27661ac08b4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round4_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a9649d6997f84508e5f566393e91e42e03a31aca55878bacd7ffd2e134bbfc +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round5_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round5_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..abd1e557c23a8769870b0beedaf7df93836a101e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round5_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7520bace4f501771d879621facd1b44f8db1ad7385a653e1535c7642465e67d +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round6_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round6_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..ebc846b7a5c587ced387cca85c4065e921c8cbe3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round6_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1951a910d1b7587459db2110e8ac43cbc906a34267372eb21d2244a900eaec46 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round7_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round7_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..5eb1370d3fa3177cb4fb710882f4b7eb706817e8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round7_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03f40a0df2437812b9e6de4e5d4e30014acb71344efa8132f19e1fc0f5925c82 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round8_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round8_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..385c599585a0973f240a027d28e985ca72d6b721 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round8_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3ef6d8247bc37b017895c9e5bbaf32f2db500f94f66b4ebade5a4a5ba7a271 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round9_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round9_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..219adb8927937a4597cd82efbf7ebb4c8a7519ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAggOnly_iter100_round5_hetero/round9_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:277fa5f0607d993e80153171e4760f999c4ee6a65452ea62694d0b9b03818f22 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..e740c469f5a7cd605c7d99d171517a9a57610ca5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c782cb176ecc64427a20748fc11e3550546d1b0313033e42d060a4d920161655 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..32dcf1d38ad238d9a1c840129bbe352efcbcf47f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3133d1a06dd7fc616ca517a5b0d8b9a7be11d1f43d25c1d70df92e151e36512 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e2a246d8d4e2073a00e90d0850f007edeed6641f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c03f477e29d8c8182396ad1132a0027ec637fe68e83220a95cb2491f921e66b9 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..3fb16e7c008686d4297b7d8f5d9c41727e2e5110 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d66a492baba1b2cf6df1e8c67d8bc995af3f2ce5447b87aef67381c198a2dff +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..433a8f6cb0f7c764d94ca00e067dfee19ee99e1e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8b3ac13d16ba28627bce9f230e91bf6e7d481fd034d6db73c84a6e93b2f28e7f +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..158d8dfab7efaa5c6d9ee3e1919776af8f7cd0b4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4589941bad011239038c28a3eec8aa55b0655bb5d0b24680d3bcd7fb72f05b51 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b0ed2187fd27d177d3bb228bcaf0bba59f22b099 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4b31ccad0fb44e5c9d811de3cfdc579ad1a2daeee509fc0e77496631805cdb3a +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..868e661c3ae337f547817510b52be75759eeea62 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:95427e7c959e2a23c257ca1a11aa61a0894e095b87f9ca5be7e285c8d78f83c8 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..4755644ac9affd0db11fba88e42695d24dbdd1b9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/0_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.2147650718688965, + "learning_rate": 2e-05, + "loss": 0.2067, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.9034323692321777, + "learning_rate": 2e-05, + "loss": 0.1734, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.3203998804092407, + "learning_rate": 2e-05, + "loss": 0.0306, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.21191878616809845, + "learning_rate": 2e-05, + "loss": 0.0434, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.71765398979187, + "learning_rate": 2e-05, + "loss": 0.2558, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.7306692600250244, + "learning_rate": 2e-05, + "loss": 0.2598, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.8116928935050964, + "learning_rate": 2e-05, + "loss": 0.0916, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.6881818771362305, + "learning_rate": 2e-05, + "loss": 0.481, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 4.576663017272949, + "learning_rate": 2e-05, + "loss": 1.0186, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.3517942428588867, + "learning_rate": 2e-05, + "loss": 0.3035, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.09975835680961609, + "learning_rate": 2e-05, + "loss": 0.0111, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.592285394668579, + "learning_rate": 2e-05, + "loss": 0.3533, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.18575400114059448, + "learning_rate": 2e-05, + "loss": 0.1854, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.2312103509902954, + "learning_rate": 2e-05, + "loss": 0.2288, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.2330929040908813, + "learning_rate": 2e-05, + "loss": 0.284, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.08729592710733414, + "learning_rate": 2e-05, + "loss": 0.0188, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.3918144702911377, + "learning_rate": 2e-05, + "loss": 0.1751, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.342497825622559, + "learning_rate": 2e-05, + "loss": 0.3139, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.7952193021774292, + "learning_rate": 2e-05, + "loss": 0.2334, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.9601510167121887, + "learning_rate": 2e-05, + "loss": 0.095, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.1995430290699005, + "learning_rate": 2e-05, + "loss": 0.0209, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.0627315044403076, + "learning_rate": 2e-05, + "loss": 0.285, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.5503566265106201, + "learning_rate": 2e-05, + "loss": 0.0972, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.3332071006298065, + "learning_rate": 2e-05, + "loss": 0.1995, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.8980509638786316, + "learning_rate": 2e-05, + "loss": 0.0735, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.6004143953323364, + "learning_rate": 2e-05, + "loss": 0.0846, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.2937055230140686, + "learning_rate": 2e-05, + "loss": 0.1981, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.7944265604019165, + "learning_rate": 2e-05, + "loss": 0.1489, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 3.6909844875335693, + "learning_rate": 2e-05, + "loss": 0.6169, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09489976614713669, + "learning_rate": 2e-05, + "loss": 0.0062, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.585684061050415, + "learning_rate": 2e-05, + "loss": 0.1512, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.5258896946907043, + "learning_rate": 2e-05, + "loss": 0.1016, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.35649293661117554, + "learning_rate": 2e-05, + "loss": 0.016, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.9681212902069092, + "learning_rate": 2e-05, + "loss": 0.0862, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.132624387741089, + "learning_rate": 2e-05, + "loss": 0.248, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.9118846654891968, + "learning_rate": 2e-05, + "loss": 0.0795, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.589687347412109, + "learning_rate": 2e-05, + "loss": 0.9654, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.185053825378418, + "learning_rate": 2e-05, + "loss": 0.4322, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.3636786937713623, + "learning_rate": 2e-05, + "loss": 0.2576, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.4627441167831421, + "learning_rate": 2e-05, + "loss": 0.316, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.5226000547409058, + "learning_rate": 2e-05, + "loss": 0.0885, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.33413028717041, + "learning_rate": 2e-05, + "loss": 0.1956, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.42517855763435364, + "learning_rate": 2e-05, + "loss": 0.0273, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.9592937231063843, + "learning_rate": 2e-05, + "loss": 0.1236, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.5450046062469482, + "learning_rate": 2e-05, + "loss": 0.1, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.6911335587501526, + "learning_rate": 2e-05, + "loss": 0.2047, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.938748598098755, + "learning_rate": 2e-05, + "loss": 1.0679, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.1489653587341309, + "learning_rate": 2e-05, + "loss": 0.4258, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.2537198066711426, + "learning_rate": 2e-05, + "loss": 0.3366, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.28481265902519226, + "learning_rate": 2e-05, + "loss": 0.0694, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289495989583872.0, + "train_loss": 0.23573887825012207, + "train_runtime": 98.4516, + "train_samples_per_second": 4.063, + "train_steps_per_second": 1.016 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289495989583872.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..546e6c50db7b2c0233dd471575687d737af2bd3c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3230800a8e8a3581ec5060154153f2e29877572afbad61dc80b1963d82b333 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..2d6cb282d15449f72dece9fcd347cb0e1a74408e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3920ebd0388798ad6921cb6afe088f0fcfda607b3aa320b4c3040772772bb247 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..72a39747645fbde80cac36b912972db95fd05b83 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:15d1dbe7114472c32be9c54ad665ea5144dfdd955454df111ae5cf5267499c89 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..faffd15a8670a8b1e11390bc5cd29374744b345c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:406da16410dba02a1c2b8904b84131e16859320f64141fa712f1f50e65775520 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5348da1cca34cd1124365960c925ff97ff88b4d8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a4ec887fb392753e18ae8649a698ee85589b37b532db636250441f2b0122403b +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..852b86739cadc784fc77008c62e8c69d83a68d8e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:798b0c0ea82489bbd134c71bff63f224b02700456156b84aa03fc529566bf3f5 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..48f98a2cf4231dca0bf506997397ba74603e2feb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8ea46935a6d9042b258eb8b0923edd0b08ae7b868be464a016dd05a4c2241051 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..23691798ba76ae32f21dff52ab916095fcf55f0c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:53a949f4b1bbcbf8b069587bb53a8e9016f8d6f6c9c28b3a4343b17acee48b4a +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..54bcef5e07c0058ee8b754f5c72544e2750131bd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/10_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.26141095161438, + "learning_rate": 2e-05, + "loss": 0.3493, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.716315746307373, + "learning_rate": 2e-05, + "loss": 0.2174, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.458876371383667, + "learning_rate": 2e-05, + "loss": 0.1478, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.4033052921295166, + "learning_rate": 2e-05, + "loss": 0.5617, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.8347094058990479, + "learning_rate": 2e-05, + "loss": 0.2959, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.15257228910923004, + "learning_rate": 2e-05, + "loss": 0.0188, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 5.624625205993652, + "learning_rate": 2e-05, + "loss": 0.4104, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.8172123432159424, + "learning_rate": 2e-05, + "loss": 0.1742, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.1313936710357666, + "learning_rate": 2e-05, + "loss": 0.1126, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.7304859161376953, + "learning_rate": 2e-05, + "loss": 0.3048, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.032973527908325, + "learning_rate": 2e-05, + "loss": 0.2187, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.114759922027588, + "learning_rate": 2e-05, + "loss": 0.1694, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.067836284637451, + "learning_rate": 2e-05, + "loss": 0.5837, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 9.950809478759766, + "learning_rate": 2e-05, + "loss": 0.6038, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 8.954267501831055, + "learning_rate": 2e-05, + "loss": 0.5741, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.319699764251709, + "learning_rate": 2e-05, + "loss": 0.0825, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.949901819229126, + "learning_rate": 2e-05, + "loss": 0.1434, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.30593255162239075, + "learning_rate": 2e-05, + "loss": 0.2789, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.5694098472595215, + "learning_rate": 2e-05, + "loss": 0.1865, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.013272285461426, + "learning_rate": 2e-05, + "loss": 0.3774, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.49882233142852783, + "learning_rate": 2e-05, + "loss": 0.0766, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.20975802838802338, + "learning_rate": 2e-05, + "loss": 0.6641, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.47412487864494324, + "learning_rate": 2e-05, + "loss": 0.0552, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.8216009736061096, + "learning_rate": 2e-05, + "loss": 0.3113, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.9364389181137085, + "learning_rate": 2e-05, + "loss": 0.1123, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.7828830480575562, + "learning_rate": 2e-05, + "loss": 0.9373, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 8.558095932006836, + "learning_rate": 2e-05, + "loss": 0.936, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.43436041474342346, + "learning_rate": 2e-05, + "loss": 0.0302, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 5.240627765655518, + "learning_rate": 2e-05, + "loss": 0.1205, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 6.766542434692383, + "learning_rate": 2e-05, + "loss": 0.4233, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 5.462146282196045, + "learning_rate": 2e-05, + "loss": 0.4414, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.6197680234909058, + "learning_rate": 2e-05, + "loss": 0.1268, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.404405355453491, + "learning_rate": 2e-05, + "loss": 0.7416, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.1363625526428223, + "learning_rate": 2e-05, + "loss": 0.3722, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.3133790493011475, + "learning_rate": 2e-05, + "loss": 0.0578, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 13.155050277709961, + "learning_rate": 2e-05, + "loss": 0.8704, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.337088584899902, + "learning_rate": 2e-05, + "loss": 0.3122, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 9.294634819030762, + "learning_rate": 2e-05, + "loss": 0.3817, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 4.370326995849609, + "learning_rate": 2e-05, + "loss": 0.7734, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.420010805130005, + "learning_rate": 2e-05, + "loss": 0.6593, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 5.748915672302246, + "learning_rate": 2e-05, + "loss": 0.3661, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.06595393270254135, + "learning_rate": 2e-05, + "loss": 0.0061, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 6.09881067276001, + "learning_rate": 2e-05, + "loss": 0.4189, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 4.024842739105225, + "learning_rate": 2e-05, + "loss": 0.5383, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.143564224243164, + "learning_rate": 2e-05, + "loss": 0.1619, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.688160419464111, + "learning_rate": 2e-05, + "loss": 0.5001, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.550379991531372, + "learning_rate": 2e-05, + "loss": 0.6729, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.03562068939209, + "learning_rate": 2e-05, + "loss": 0.1165, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.7042808532714844, + "learning_rate": 2e-05, + "loss": 0.3509, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.446930468082428, + "learning_rate": 2e-05, + "loss": 0.4792, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2220576018006016.0, + "train_loss": 0.3565188694000244, + "train_runtime": 60.1393, + "train_samples_per_second": 6.651, + "train_steps_per_second": 1.663 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2220576018006016.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..37384d9fb9305423c29b82651387afa2f469df7d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5a551340d5ac72a4337e2f038bb19911acaf327b4a6649864b1aef0639fed08d +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..8c019eac259435f0b203a96e681defb93ada3927 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:13b3015737c54b64eac148c7069ad431c9361f991267f0e2dba6f4ec1e194e5c +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..984e845b607e8a2dcfcc5e6c96e9ebc8c7b122c9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d313f7a70916a5b629c99f5b47164da2eaefce5e9b7f3832f3de43fcd4f18814 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b48305418651df12eceac332ef53db184b20b789 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:74851bce7f8704d25a007f89268a27bfb39cc2415b70fe4d7536b0b8902dff5d +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ae477a002cda808b2b2ed6428c422b8873064a2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5a6fced4b804c6acb6d815ba6982d27fa3f5a9ca7deca288efad1ae69fc84b5d +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..f03aa8c68d0d402458526182e30b15ba70fe0008 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b0b110252f8d3f03a1f76860e2927c475e79763013bd88edd6414bd8f954f5e8 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..afc7a4315d72b64ea66fa2e8538bd770c716bf23 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:27a0aa78bf886df4f04676769bdd911b4af61f0a9fc036c9b9dc05c4f3eaa1cc +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..e8ffc49c4bce8beffe545e26c74d2d774818dd4f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1a53fa9c6bae2940016edc9c0e898acce0d49f7e9749e101a7aaa028218181d7 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..bc7b81b7f3b36045a59f3ab8523dbe40f8031a4c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/11_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.3963344097137451, + "learning_rate": 2e-05, + "loss": 0.5264, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.1859383583068848, + "learning_rate": 2e-05, + "loss": 0.4709, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.9786797761917114, + "learning_rate": 2e-05, + "loss": 0.3583, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.240707278251648, + "learning_rate": 2e-05, + "loss": 0.3753, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.249172687530518, + "learning_rate": 2e-05, + "loss": 0.752, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.8528921008110046, + "learning_rate": 2e-05, + "loss": 0.3783, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.4588570594787598, + "learning_rate": 2e-05, + "loss": 0.874, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.918217897415161, + "learning_rate": 2e-05, + "loss": 0.7202, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.7528201341629028, + "learning_rate": 2e-05, + "loss": 0.5347, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.7792121171951294, + "learning_rate": 2e-05, + "loss": 0.3746, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.5276216864585876, + "learning_rate": 2e-05, + "loss": 0.2802, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.8013849854469299, + "learning_rate": 2e-05, + "loss": 0.4915, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.5848059058189392, + "learning_rate": 2e-05, + "loss": 0.3829, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.433096170425415, + "learning_rate": 2e-05, + "loss": 0.3877, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.8524357676506042, + "learning_rate": 2e-05, + "loss": 0.4814, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.8882460594177246, + "learning_rate": 2e-05, + "loss": 0.407, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.0338265895843506, + "learning_rate": 2e-05, + "loss": 0.4946, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.0775980949401855, + "learning_rate": 2e-05, + "loss": 0.4995, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.5330630540847778, + "learning_rate": 2e-05, + "loss": 0.5254, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.5867397785186768, + "learning_rate": 2e-05, + "loss": 0.45, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.3602073192596436, + "learning_rate": 2e-05, + "loss": 0.3577, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.2039196491241455, + "learning_rate": 2e-05, + "loss": 0.5869, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.7619243860244751, + "learning_rate": 2e-05, + "loss": 0.4807, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.41163066029548645, + "learning_rate": 2e-05, + "loss": 0.2659, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.8487000465393066, + "learning_rate": 2e-05, + "loss": 0.6214, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.184577703475952, + "learning_rate": 2e-05, + "loss": 0.4849, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.9186478853225708, + "learning_rate": 2e-05, + "loss": 0.3281, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.9017455577850342, + "learning_rate": 2e-05, + "loss": 0.2539, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.4897398948669434, + "learning_rate": 2e-05, + "loss": 0.2508, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5702741146087646, + "learning_rate": 2e-05, + "loss": 0.3073, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.3382965326309204, + "learning_rate": 2e-05, + "loss": 0.4023, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.7662063837051392, + "learning_rate": 2e-05, + "loss": 0.3478, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.596600294113159, + "learning_rate": 2e-05, + "loss": 0.6128, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.08713281899690628, + "learning_rate": 2e-05, + "loss": 0.1682, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.2343882322311401, + "learning_rate": 2e-05, + "loss": 0.3383, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.757921576499939, + "learning_rate": 2e-05, + "loss": 0.4687, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.496662974357605, + "learning_rate": 2e-05, + "loss": 0.665, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.7695884704589844, + "learning_rate": 2e-05, + "loss": 0.3334, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 3.126971483230591, + "learning_rate": 2e-05, + "loss": 0.5775, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.082252025604248, + "learning_rate": 2e-05, + "loss": 0.6768, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.217541217803955, + "learning_rate": 2e-05, + "loss": 0.7842, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.6298086643218994, + "learning_rate": 2e-05, + "loss": 0.6875, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.039024591445923, + "learning_rate": 2e-05, + "loss": 0.2957, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.9569606781005859, + "learning_rate": 2e-05, + "loss": 0.4915, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.017443710938096046, + "learning_rate": 2e-05, + "loss": 0.5058, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.0939208269119263, + "learning_rate": 2e-05, + "loss": 0.3516, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.7852720618247986, + "learning_rate": 2e-05, + "loss": 0.4697, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.0443236827850342, + "learning_rate": 2e-05, + "loss": 0.439, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.216949462890625, + "learning_rate": 2e-05, + "loss": 0.4072, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.3386884927749634, + "learning_rate": 2e-05, + "loss": 0.4341, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2191347477905408.0, + "train_loss": 0.4631907272338867, + "train_runtime": 78.9468, + "train_samples_per_second": 5.067, + "train_steps_per_second": 1.267 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2191347477905408.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..23fbac9a2dd4fabaf9090e1da38215212a26534c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3271ba524ae51a79308b2b0ed607d9581f14a67b55dc1a0b49e9050b77c4afa8 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ed95b1a720cc81a26ebc9ba53d5c0fab0e623363 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d905ae63bfaf963d27434307670f258dd6ca9fb2875fde22eaa7689148a27af5 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a9670fd6e8d7ddd753036e82880e328fd90c3ce7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:294bcd98746724a2062c77391edf700678c08ebfa2101bf7c5c84a779834c41d +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..c43a6bc552b2c5f155ab31327aab06fb6ff65267 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1f580d8eada4621e38ddf406c915f0835cf7b035d0ccfdeee9d586f3053ab304 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..3ea5ecee85cd6b5db43fe4524e1f68a980f64a64 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76667e589eb0e73875c744c238bb4eecf6a129ee382d032187d0a97ffec0c226 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..277168ec4afa1a84e8e8078ab2c21bcde9181d4b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6d1b63d6760cf9a278ff19626e7ccfb02514b796d51d29bf37904f5b88318cc0 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..27822a4f69a78069fde7644573e5cc8aa9143001 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6278768c22e24626aa97b9f1a7d947d6c62b6ef5b3bd8b26e2af9a65a7f5dae4 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..efb7c6bcbf21d1aa59976f42aad14fd406f6379b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3ef6b01dde335f74273cca747ee62ff48b2b14c7b8c80014ebba98076f7f2e6a +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..a8d617775a8b4d8a063f259503f2b8c37255e0a6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/12_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.8910015225410461, + "learning_rate": 2e-05, + "loss": 0.0391, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.0014869185397401452, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.06336291134357452, + "learning_rate": 2e-05, + "loss": 0.0136, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.5776440501213074, + "learning_rate": 2e-05, + "loss": 0.0291, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.34741267561912537, + "learning_rate": 2e-05, + "loss": 0.0201, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.056197356432676315, + "learning_rate": 2e-05, + "loss": 0.0025, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.9980051517486572, + "learning_rate": 2e-05, + "loss": 0.5898, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.9883863925933838, + "learning_rate": 2e-05, + "loss": 0.2791, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.2480560690164566, + "learning_rate": 2e-05, + "loss": 0.167, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.017595052719116, + "learning_rate": 2e-05, + "loss": 0.1544, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.1280212104320526, + "learning_rate": 2e-05, + "loss": 0.1391, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.15525102615356445, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.22820474207401276, + "learning_rate": 2e-05, + "loss": 0.0493, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.5431236028671265, + "learning_rate": 2e-05, + "loss": 0.1274, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.02224026434123516, + "learning_rate": 2e-05, + "loss": 0.0104, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.042856115847826004, + "learning_rate": 2e-05, + "loss": 0.0118, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.08032882213592529, + "learning_rate": 2e-05, + "loss": 0.0332, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.11706147342920303, + "learning_rate": 2e-05, + "loss": 0.0795, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.15025536715984344, + "learning_rate": 2e-05, + "loss": 0.0244, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.4260348379611969, + "learning_rate": 2e-05, + "loss": 0.0363, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.3484039306640625, + "learning_rate": 2e-05, + "loss": 0.3069, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.018344158306717873, + "learning_rate": 2e-05, + "loss": 0.0245, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.0765979290008545, + "learning_rate": 2e-05, + "loss": 0.1697, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.5943951606750488, + "learning_rate": 2e-05, + "loss": 0.0322, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.06011451780796051, + "learning_rate": 2e-05, + "loss": 0.0029, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.08369926363229752, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.03558244928717613, + "learning_rate": 2e-05, + "loss": 0.0017, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.07649137824773788, + "learning_rate": 2e-05, + "loss": 0.0031, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.01627441868185997, + "learning_rate": 2e-05, + "loss": 0.0026, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.06631405651569366, + "learning_rate": 2e-05, + "loss": 0.0041, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.008456850424408913, + "learning_rate": 2e-05, + "loss": 0.0072, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.009240486659109592, + "learning_rate": 2e-05, + "loss": 0.0056, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.08000624179840088, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.5483933091163635, + "learning_rate": 2e-05, + "loss": 0.0269, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.003977715503424406, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.026056179776787758, + "learning_rate": 2e-05, + "loss": 0.2847, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.029333729296922684, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.007162914145737886, + "learning_rate": 2e-05, + "loss": 0.0178, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.373212993144989, + "learning_rate": 2e-05, + "loss": 0.1932, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.004001216031610966, + "learning_rate": 2e-05, + "loss": 0.0065, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.03465154394507408, + "learning_rate": 2e-05, + "loss": 0.0021, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.03863660618662834, + "learning_rate": 2e-05, + "loss": 0.0043, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.00269978865981102, + "learning_rate": 2e-05, + "loss": 0.0077, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.5809215307235718, + "learning_rate": 2e-05, + "loss": 0.1792, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.0029940425883978605, + "learning_rate": 2e-05, + "loss": 0.0119, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.3479861915111542, + "learning_rate": 2e-05, + "loss": 0.0733, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.2719109058380127, + "learning_rate": 2e-05, + "loss": 0.3593, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.02369946427643299, + "learning_rate": 2e-05, + "loss": 0.4603, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.022478559985756874, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.2390679121017456, + "learning_rate": 2e-05, + "loss": 0.0379, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5285302658662400.0, + "train_loss": 0.08107412010431289, + "train_runtime": 99.8411, + "train_samples_per_second": 4.006, + "train_steps_per_second": 1.002 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5285302658662400.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..f4dc06e5e6bd3bb4c4e28edd5c08256f70d29c28 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8daec27eaa622b84b1f0ef4c3a6887decfc27768f4ce882a24e8d9ea69f30f82 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..1fbf38fb88c5e3f6b01fc8e9e7b74ac4d87e535d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:731ded1d1dab877170e7994f88237e72d91a2246af5d5b8e25ed7c7ae18b947b +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..f3463f4cb942c0ae210f6af9960bc4191e10f080 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:715388a10f41bb4a8f5d79b8f09748241daa97621c3d891a203e16c3d178d36b +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..242039042283798a0e7d0e0e0be69356760e132e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6e3db23c2bf2ee6eb9b96bf40e4a45de6b112fbf02a229d20d47da7ac7b94658 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..a92e423427b51d18b0020287bb97054b08c10f34 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d80256f0abdf185fe06548e8739a276bbded6fa21c45fd616f82098728d8cb19 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..317766fcefa0257de7b09dd465a59f2750d894f7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1f3c4c9da717b3a7672770d5080193d72f0a4a2ff0a67dba4dde47ccab28f55a +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..5c9d84e87fe9cd6a5b316a5a2dd4efbd271b7c6f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ca8f694133a32e07a82f14dcb4f1231c15761ed8591cb2e708d92e215ce87345 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..63e50d472eaa8856ef48a63aef3b2aae66f02e16 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d1790f05f442700038098acb705764ed36d8dec5935a839cb2c1db61f57ba00 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..67a389b41cc3807afec4e85b7e9a0be77e4e8e4a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/13_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.9386892318725586, + "learning_rate": 2e-05, + "loss": 0.2984, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.4930660724639893, + "learning_rate": 2e-05, + "loss": 0.2379, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.656134843826294, + "learning_rate": 2e-05, + "loss": 0.3763, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.9849265217781067, + "learning_rate": 2e-05, + "loss": 0.0869, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.6041097044944763, + "learning_rate": 2e-05, + "loss": 0.0922, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.1545434594154358, + "learning_rate": 2e-05, + "loss": 0.0313, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.8947466611862183, + "learning_rate": 2e-05, + "loss": 0.0989, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.9009625911712646, + "learning_rate": 2e-05, + "loss": 0.42, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.6293833255767822, + "learning_rate": 2e-05, + "loss": 0.192, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.9993780851364136, + "learning_rate": 2e-05, + "loss": 1.233, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.870278358459473, + "learning_rate": 2e-05, + "loss": 0.5415, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 5.783886909484863, + "learning_rate": 2e-05, + "loss": 0.2204, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.167943239212036, + "learning_rate": 2e-05, + "loss": 0.9906, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.5522267818450928, + "learning_rate": 2e-05, + "loss": 0.1408, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.4071619510650635, + "learning_rate": 2e-05, + "loss": 0.3004, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.1383166313171387, + "learning_rate": 2e-05, + "loss": 0.124, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.7739863395690918, + "learning_rate": 2e-05, + "loss": 0.205, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.684206247329712, + "learning_rate": 2e-05, + "loss": 0.3239, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.448175072669983, + "learning_rate": 2e-05, + "loss": 0.1222, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.548603892326355, + "learning_rate": 2e-05, + "loss": 0.2351, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.03315643221139908, + "learning_rate": 2e-05, + "loss": 0.2654, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.543233871459961, + "learning_rate": 2e-05, + "loss": 0.4289, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.01659698225557804, + "learning_rate": 2e-05, + "loss": 0.0139, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.9258449077606201, + "learning_rate": 2e-05, + "loss": 0.2115, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.3092912435531616, + "learning_rate": 2e-05, + "loss": 0.2105, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.2467241287231445, + "learning_rate": 2e-05, + "loss": 0.427, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.9891356229782104, + "learning_rate": 2e-05, + "loss": 0.3654, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.84932941198349, + "learning_rate": 2e-05, + "loss": 0.194, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.0737885981798172, + "learning_rate": 2e-05, + "loss": 0.034, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.10633374005556107, + "learning_rate": 2e-05, + "loss": 0.0646, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.385871171951294, + "learning_rate": 2e-05, + "loss": 0.178, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.35884788632392883, + "learning_rate": 2e-05, + "loss": 0.0777, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.087784767150879, + "learning_rate": 2e-05, + "loss": 0.3361, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.607100248336792, + "learning_rate": 2e-05, + "loss": 0.2921, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.2213834524154663, + "learning_rate": 2e-05, + "loss": 0.2264, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.28556424379348755, + "learning_rate": 2e-05, + "loss": 0.0339, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5748167037963867, + "learning_rate": 2e-05, + "loss": 0.5091, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.6310743689537048, + "learning_rate": 2e-05, + "loss": 0.1153, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.39888334274292, + "learning_rate": 2e-05, + "loss": 0.2533, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.0574818849563599, + "learning_rate": 2e-05, + "loss": 0.0827, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.6757338643074036, + "learning_rate": 2e-05, + "loss": 0.1152, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 4.007380485534668, + "learning_rate": 2e-05, + "loss": 0.7782, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.645462989807129, + "learning_rate": 2e-05, + "loss": 0.5895, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.1675024032592773, + "learning_rate": 2e-05, + "loss": 0.1357, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.1133456230163574, + "learning_rate": 2e-05, + "loss": 0.1158, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.6135990619659424, + "learning_rate": 2e-05, + "loss": 0.5882, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.10320578515529633, + "learning_rate": 2e-05, + "loss": 0.0105, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.8385117053985596, + "learning_rate": 2e-05, + "loss": 0.2106, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.34526586532592773, + "learning_rate": 2e-05, + "loss": 0.3233, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.3280060291290283, + "learning_rate": 2e-05, + "loss": 0.2871, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5322110830379008.0, + "train_loss": 0.2748959970474243, + "train_runtime": 92.4298, + "train_samples_per_second": 4.328, + "train_steps_per_second": 1.082 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5322110830379008.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..720beb032170c365d0c3cb0f45553cee5e58d3fa --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ae603a400db719c331df7011557f624adb310ab0c5671b72b7c6bcb6c17a40fa +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..ecf102eb0dfee8003a345d4c7250c7f3fc16b4e6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c4256c3a83725cdf202948e2af7b9e08ca0d2cb1703e542c12601a4999e76b4 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e5817202b3eec1cd4ce25a39e35ff58e0c57c800 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4bdc426c17379e947c0a82b43e2671c71e48fa71f0992dd5afd4f1f8a67e239f +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..052804230e0ea52690f1be45be6699ed620266a4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8359e73f5f55ceddf5a7084c5be15738aa2ac846f394a451ceefda9404b95ce3 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..bf1263f90f799e6091958eb31a0b52ec3894960e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4cbf39dfc598ed79b6f990541ad236921ace95406ead69971280a7116e26e3de +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..9eeb2a9fc865b94d96cdcd76d28270c87d1aa243 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aae4da68fd2ab6c44f72f304dcf4311bb6f2c75b75de60f26b0bd7cc8e1de31b +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..50c06afae0554f66ad5aae6a5d19ebb96a6e05ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c0146cbffd991cdf93235c41285d847df280ee015a325dce05f55b50ddaaad87 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..3de0f32e0a4f48fc945ebdc115e0734e2b096582 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4d3b8778d52137c0fda850e50dcf94b6ef9cae707497cb4f0f372294c6519ace +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..0b53b4a87d8952757fa9392750e0ce05411df8cf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/14_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.5939472317695618, + "learning_rate": 2e-05, + "loss": 0.4598, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.6588802337646484, + "learning_rate": 2e-05, + "loss": 0.2883, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.1761450171470642, + "learning_rate": 2e-05, + "loss": 0.0502, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.05561560019850731, + "learning_rate": 2e-05, + "loss": 0.0014, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 5.902976036071777, + "learning_rate": 2e-05, + "loss": 0.6313, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.9888331294059753, + "learning_rate": 2e-05, + "loss": 0.2214, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.900629997253418, + "learning_rate": 2e-05, + "loss": 0.3104, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.023601477965712547, + "learning_rate": 2e-05, + "loss": 0.0032, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.0727896690368652, + "learning_rate": 2e-05, + "loss": 0.1358, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.30931559205055237, + "learning_rate": 2e-05, + "loss": 0.0252, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.13978511095046997, + "learning_rate": 2e-05, + "loss": 0.024, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.03559078276157379, + "learning_rate": 2e-05, + "loss": 0.0021, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 5.685776710510254, + "learning_rate": 2e-05, + "loss": 0.614, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 6.392800331115723, + "learning_rate": 2e-05, + "loss": 0.3703, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.114696502685547, + "learning_rate": 2e-05, + "loss": 0.2291, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.9133680462837219, + "learning_rate": 2e-05, + "loss": 0.0802, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.15922507643699646, + "learning_rate": 2e-05, + "loss": 0.2088, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.642828643321991, + "learning_rate": 2e-05, + "loss": 0.0398, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.09540918469429016, + "learning_rate": 2e-05, + "loss": 0.0132, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.401083469390869, + "learning_rate": 2e-05, + "loss": 0.3875, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.25489044189453125, + "learning_rate": 2e-05, + "loss": 0.1419, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.5208061933517456, + "learning_rate": 2e-05, + "loss": 0.3872, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.2758226990699768, + "learning_rate": 2e-05, + "loss": 0.0325, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.7475370168685913, + "learning_rate": 2e-05, + "loss": 0.047, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.02500426024198532, + "learning_rate": 2e-05, + "loss": 0.184, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.03836209326982498, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.5272471904754639, + "learning_rate": 2e-05, + "loss": 0.0535, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.7803357243537903, + "learning_rate": 2e-05, + "loss": 0.1147, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.820475697517395, + "learning_rate": 2e-05, + "loss": 0.1912, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.015785478055477142, + "learning_rate": 2e-05, + "loss": 0.0032, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.6323795914649963, + "learning_rate": 2e-05, + "loss": 0.2556, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.28682807087898254, + "learning_rate": 2e-05, + "loss": 0.0251, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.6010555028915405, + "learning_rate": 2e-05, + "loss": 0.106, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.629011631011963, + "learning_rate": 2e-05, + "loss": 0.1921, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.4394206702709198, + "learning_rate": 2e-05, + "loss": 0.165, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.17273245751857758, + "learning_rate": 2e-05, + "loss": 0.0321, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.7993721961975098, + "learning_rate": 2e-05, + "loss": 0.1456, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.1527528762817383, + "learning_rate": 2e-05, + "loss": 0.2189, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.0030514001846313, + "learning_rate": 2e-05, + "loss": 0.1274, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.151094675064087, + "learning_rate": 2e-05, + "loss": 0.0861, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.35782918334007263, + "learning_rate": 2e-05, + "loss": 0.045, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.6118261814117432, + "learning_rate": 2e-05, + "loss": 0.0487, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.5953887104988098, + "learning_rate": 2e-05, + "loss": 0.0228, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.5005707740783691, + "learning_rate": 2e-05, + "loss": 0.029, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.8890740871429443, + "learning_rate": 2e-05, + "loss": 0.0845, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.2526500225067139, + "learning_rate": 2e-05, + "loss": 0.034, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.10460038483142853, + "learning_rate": 2e-05, + "loss": 0.0162, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.2411116361618042, + "learning_rate": 2e-05, + "loss": 0.0284, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.2785385251045227, + "learning_rate": 2e-05, + "loss": 0.0141, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.003894544905051589, + "learning_rate": 2e-05, + "loss": 0.0354, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5337401710870528.0, + "train_loss": 0.13932286381721495, + "train_runtime": 105.3402, + "train_samples_per_second": 3.797, + "train_steps_per_second": 0.949 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5337401710870528.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..998b07d69280bb88f99b49418cd1424b8838d7ca --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:524c3ee8320b86ce59b738adb9f47d3f09ff4ac9314441a72ab2dbf3be6f7ed9 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..edad600d1f6d06fc287fa536ece938a73f033e97 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:46363767655b011ca85fc67bc88224679143dacc45a9652ff1565700c998bcc0 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..d4a511f7326caa7ecc4d74848ce24782d73d5cc2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eec4bb5d2aada2a62fa1f709fc4485a5a6feed4628be81f95ad118ec6649203e +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..4e6cf995b17f00d9a13c8becf32d7c39b9a27de6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2c1e39479b43c9201ffff766f753b35a15f19d435dddb251913fcec12deeb896 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5a27832cae52ff9453f85c8e29dc02ef6f1ded35 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0130e1d4ddae5b43948c71805445f40ea600f24e573f1faec0a7ac7b3f55388c +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..510e36cf1070edae1924b0c5f8f32fbf7c023ad5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b4f170426d7c168d474be30ba0f265b96ad230867a26e03168314ccff3e3aa1 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..32b80ceaee7679d8cf43697bcdf19e6c7004de58 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2359df50f5315d1e8480b1ad339b7f741f86699f4c3a51b64019e9de405bad0c +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..0214bb5eb16ccc4dff14a26e7db14fb054ee7c0e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ca64fecf6d1049c525e49ccd17a31270e5368230b1b24b2f7939e168fbb8c4c9 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..17eb85da83c73a3d23c899d2a19af27b5a57dd24 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/15_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.534881591796875, + "learning_rate": 2e-05, + "loss": 0.1539, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.637804985046387, + "learning_rate": 2e-05, + "loss": 0.3872, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 4.103468894958496, + "learning_rate": 2e-05, + "loss": 0.4706, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.3632436990737915, + "learning_rate": 2e-05, + "loss": 0.1057, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.46576496958732605, + "learning_rate": 2e-05, + "loss": 0.0883, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.4900715947151184, + "learning_rate": 2e-05, + "loss": 0.1953, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.6965852975845337, + "learning_rate": 2e-05, + "loss": 0.0454, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.3500645160675049, + "learning_rate": 2e-05, + "loss": 0.2741, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.987757682800293, + "learning_rate": 2e-05, + "loss": 0.1152, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.1680185794830322, + "learning_rate": 2e-05, + "loss": 0.1962, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.889827728271484, + "learning_rate": 2e-05, + "loss": 0.4485, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.615948438644409, + "learning_rate": 2e-05, + "loss": 0.203, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.9771978855133057, + "learning_rate": 2e-05, + "loss": 0.3898, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.4878509044647217, + "learning_rate": 2e-05, + "loss": 0.2237, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.5994541645050049, + "learning_rate": 2e-05, + "loss": 0.2438, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.6307693719863892, + "learning_rate": 2e-05, + "loss": 0.0721, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.4535841941833496, + "learning_rate": 2e-05, + "loss": 0.4117, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.1207384467124939, + "learning_rate": 2e-05, + "loss": 0.0447, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.8656939268112183, + "learning_rate": 2e-05, + "loss": 0.2773, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.4517691135406494, + "learning_rate": 2e-05, + "loss": 0.1175, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.20167458057403564, + "learning_rate": 2e-05, + "loss": 0.0954, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.9522814750671387, + "learning_rate": 2e-05, + "loss": 0.5457, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.3511316776275635, + "learning_rate": 2e-05, + "loss": 0.038, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.421407699584961, + "learning_rate": 2e-05, + "loss": 0.0943, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.124981880187988, + "learning_rate": 2e-05, + "loss": 0.6406, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.5482290983200073, + "learning_rate": 2e-05, + "loss": 0.2579, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.12442682683467865, + "learning_rate": 2e-05, + "loss": 0.0166, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4804109036922455, + "learning_rate": 2e-05, + "loss": 0.0449, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.9573951959609985, + "learning_rate": 2e-05, + "loss": 0.2828, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.6679618954658508, + "learning_rate": 2e-05, + "loss": 0.0887, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 8.849357604980469, + "learning_rate": 2e-05, + "loss": 0.6207, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.9769513607025146, + "learning_rate": 2e-05, + "loss": 0.1143, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.2505854368209839, + "learning_rate": 2e-05, + "loss": 0.1845, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.904111385345459, + "learning_rate": 2e-05, + "loss": 0.5034, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.4404342770576477, + "learning_rate": 2e-05, + "loss": 0.3872, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.3586173057556152, + "learning_rate": 2e-05, + "loss": 0.2092, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.4986414909362793, + "learning_rate": 2e-05, + "loss": 0.1584, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.6520402431488037, + "learning_rate": 2e-05, + "loss": 0.3777, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.9505988359451294, + "learning_rate": 2e-05, + "loss": 0.0464, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.3538886308670044, + "learning_rate": 2e-05, + "loss": 0.1724, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.5838003158569336, + "learning_rate": 2e-05, + "loss": 0.1721, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.1229430437088013, + "learning_rate": 2e-05, + "loss": 0.1696, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.624340772628784, + "learning_rate": 2e-05, + "loss": 0.542, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 5.232175827026367, + "learning_rate": 2e-05, + "loss": 0.6533, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.7528300285339355, + "learning_rate": 2e-05, + "loss": 0.2582, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.3056037724018097, + "learning_rate": 2e-05, + "loss": 0.1694, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.7677148580551147, + "learning_rate": 2e-05, + "loss": 0.0556, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.5587077140808105, + "learning_rate": 2e-05, + "loss": 0.7456, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.238245725631714, + "learning_rate": 2e-05, + "loss": 0.2669, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.7615108489990234, + "learning_rate": 2e-05, + "loss": 0.1252, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2213537237696512.0, + "train_loss": 0.2500221920013428, + "train_runtime": 62.3198, + "train_samples_per_second": 6.419, + "train_steps_per_second": 1.605 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2213537237696512.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..b6bf2f0b03ee017b2bbf430f9e3f8c1e6a7ce5c8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:432c39ecf5095d9ba971af1e6980b8dc47a4e1ef37b9b3b25716e93de8d0d498 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..9eaf5a2b38cbcfd23a93db97225423090193961c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7527eb6d0a5f947a0b644ff985fac521b9232d5fe684817c22d8fce96fcfa07a +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..30834b8e33528a9c9f1ce796aeb21baf0fd22433 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:55d31af1c7c7f04705f71a76c1a048e2be7ab4191a01d51200fbc46c31a95bad +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..ae2933eb7954afbf36e93a51f4497ee6d6b0542e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ef10f82c3555a4adbd4717faf8119d4bccafeddf29c6f9416e6f6b2541c5f430 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..8cdc3f5e0085ef4a2632169379575ed65edbe087 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:781c8155f3b5c351590e41e2a15ea770802b3413e3bf68b4886b12083f9f3808 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..1085d56c27504f0a16d6087d50c96a4166853904 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be084abc13154c1117982406f174710f1878ebd1f04f99eda443c15cf7e83da5 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..6bf1a0ff77164dcc2572bb9827cc28292505223e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc0d8a7a404470068cf2a7544b6221a1605d34ec46e1981c460badf1e62643f3 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..10eac913a2807a2f40d1acdd6825f5cd0a6d8ef5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29d1b63b674861b08d9b2d3ae2f9e5afb18f98bfe11a78f87a37d07fdebf4041 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..96908fa90cb248b467705d0f850e1d9909482280 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/16_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.0506765842437744, + "learning_rate": 2e-05, + "loss": 0.2656, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.11302793771028519, + "learning_rate": 2e-05, + "loss": 0.0321, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 8.295666694641113, + "learning_rate": 2e-05, + "loss": 0.8303, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 4.383336067199707, + "learning_rate": 2e-05, + "loss": 0.2309, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.513815402984619, + "learning_rate": 2e-05, + "loss": 0.2053, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.09286098927259445, + "learning_rate": 2e-05, + "loss": 0.025, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.9971388578414917, + "learning_rate": 2e-05, + "loss": 0.1517, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.6069109439849854, + "learning_rate": 2e-05, + "loss": 0.3094, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.362154483795166, + "learning_rate": 2e-05, + "loss": 0.0705, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.08220279961824417, + "learning_rate": 2e-05, + "loss": 0.0158, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.907205581665039, + "learning_rate": 2e-05, + "loss": 0.6874, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.4044019877910614, + "learning_rate": 2e-05, + "loss": 0.1621, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.035224199295044, + "learning_rate": 2e-05, + "loss": 0.4915, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.4353359937667847, + "learning_rate": 2e-05, + "loss": 0.1595, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.4151556491851807, + "learning_rate": 2e-05, + "loss": 0.5463, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.7523651123046875, + "learning_rate": 2e-05, + "loss": 0.245, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.2892554700374603, + "learning_rate": 2e-05, + "loss": 0.2165, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.6359217166900635, + "learning_rate": 2e-05, + "loss": 0.2207, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.375725507736206, + "learning_rate": 2e-05, + "loss": 0.0974, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.7544479370117188, + "learning_rate": 2e-05, + "loss": 0.2783, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.08416328579187393, + "learning_rate": 2e-05, + "loss": 0.0068, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.5543112754821777, + "learning_rate": 2e-05, + "loss": 0.1415, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.1626098155975342, + "learning_rate": 2e-05, + "loss": 0.0876, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.3678817749023438, + "learning_rate": 2e-05, + "loss": 0.0911, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.1967419534921646, + "learning_rate": 2e-05, + "loss": 0.0548, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.5340510606765747, + "learning_rate": 2e-05, + "loss": 0.0467, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.7714657187461853, + "learning_rate": 2e-05, + "loss": 0.0432, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 5.252527713775635, + "learning_rate": 2e-05, + "loss": 0.403, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.5596518516540527, + "learning_rate": 2e-05, + "loss": 0.2485, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5285986661911011, + "learning_rate": 2e-05, + "loss": 0.0576, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.3168255090713501, + "learning_rate": 2e-05, + "loss": 0.0151, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.5153025388717651, + "learning_rate": 2e-05, + "loss": 0.5059, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.5826822519302368, + "learning_rate": 2e-05, + "loss": 0.0978, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 8.49281120300293, + "learning_rate": 2e-05, + "loss": 1.6427, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.3995786905288696, + "learning_rate": 2e-05, + "loss": 0.1928, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 7.620669364929199, + "learning_rate": 2e-05, + "loss": 0.3542, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.6757049560546875, + "learning_rate": 2e-05, + "loss": 0.7299, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.033230576664209366, + "learning_rate": 2e-05, + "loss": 0.0032, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.364359736442566, + "learning_rate": 2e-05, + "loss": 0.0363, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.26916366815567017, + "learning_rate": 2e-05, + "loss": 0.0299, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.4223482608795166, + "learning_rate": 2e-05, + "loss": 0.4501, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.9469695091247559, + "learning_rate": 2e-05, + "loss": 0.2121, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.2505552768707275, + "learning_rate": 2e-05, + "loss": 0.0329, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.1719236373901367, + "learning_rate": 2e-05, + "loss": 0.2281, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.545175552368164, + "learning_rate": 2e-05, + "loss": 0.4012, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.245997667312622, + "learning_rate": 2e-05, + "loss": 0.1146, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.817222595214844, + "learning_rate": 2e-05, + "loss": 0.7762, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 5.109766960144043, + "learning_rate": 2e-05, + "loss": 0.1387, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.6398802995681763, + "learning_rate": 2e-05, + "loss": 0.104, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.171039581298828, + "learning_rate": 2e-05, + "loss": 0.1401, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206762417520640.0, + "train_loss": 0.25256196975708006, + "train_runtime": 62.66, + "train_samples_per_second": 6.384, + "train_steps_per_second": 1.596 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206762417520640.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..97be9c6052c6c62d660d28564eb0eefd2b059367 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4b4494f4a6c0fd73f702b17bef5c18506747118dbf76374343bbc66ac493c9dc +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..3ffe17c093a81e85312c717af2e2588942b782a7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b1daa9eb5a30be63a3015a04e61a0641a10ce96b49109d8f4d39e48984dd91e6 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..c174de36be52b70ede1dd44958c8ec5ab7680c8f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:14b9e8601bf2c03d31b3a728f286d007965d0796ca2b360cd6d7399d81ff568b +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..6fd267e1b27139c485feb299c719e37b86a419c8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2d9b13d4fbd249fb850d1bfce848122b6ad8031ab5f4a7badd01e5aafa034306 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..0ac5a831557bff0f857f2585d7a42347f9b352b9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:82e30397000e603146575d0ab1f94053c8cacaabb759369f44971b7d0e1b34a8 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..d5cdf88709b565af5293873b8ffac34bf37e447b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29605b0a429e4afe9ba599d70bd16f89b31b2979d18eafbeea995fa7312e69aa +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..eb84ee7a91bfad6a5e10a5795c17c2120b096761 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8490f5799b1408d689e89f1c4a347fc6dabd4c15dcafbff28b7da1bddfe02d43 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..577b2e644d275c3f9cbc770fda97da94394a94cf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1ab3a258077bb8851ecd26276ba54d935f617f2c960ae768668394d722f64d67 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..9ca260a8cfc9b362d71e8832d900ea99f47fea37 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/17_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.5502195358276367, + "learning_rate": 2e-05, + "loss": 0.0507, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.9286928176879883, + "learning_rate": 2e-05, + "loss": 0.4337, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.9705684185028076, + "learning_rate": 2e-05, + "loss": 0.113, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.7228869795799255, + "learning_rate": 2e-05, + "loss": 0.0544, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.5402493476867676, + "learning_rate": 2e-05, + "loss": 0.1068, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.2117934226989746, + "learning_rate": 2e-05, + "loss": 0.1799, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.6345689296722412, + "learning_rate": 2e-05, + "loss": 0.0237, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.215928554534912, + "learning_rate": 2e-05, + "loss": 0.2292, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 5.735645294189453, + "learning_rate": 2e-05, + "loss": 0.3665, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 6.8621296882629395, + "learning_rate": 2e-05, + "loss": 0.9288, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.483808755874634, + "learning_rate": 2e-05, + "loss": 0.2017, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 5.605947971343994, + "learning_rate": 2e-05, + "loss": 0.6995, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.7000005841255188, + "learning_rate": 2e-05, + "loss": 0.0806, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.12534011900424957, + "learning_rate": 2e-05, + "loss": 0.0089, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.2789242267608643, + "learning_rate": 2e-05, + "loss": 0.1396, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.786081075668335, + "learning_rate": 2e-05, + "loss": 0.1208, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8238869905471802, + "learning_rate": 2e-05, + "loss": 0.0716, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3384992182254791, + "learning_rate": 2e-05, + "loss": 0.019, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.3042421340942383, + "learning_rate": 2e-05, + "loss": 0.2698, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.305187702178955, + "learning_rate": 2e-05, + "loss": 0.2792, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.3464645743370056, + "learning_rate": 2e-05, + "loss": 0.0253, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.6517714262008667, + "learning_rate": 2e-05, + "loss": 0.2873, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.8206135034561157, + "learning_rate": 2e-05, + "loss": 0.0613, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.6411991119384766, + "learning_rate": 2e-05, + "loss": 0.2191, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.8292698860168457, + "learning_rate": 2e-05, + "loss": 0.3554, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.263585090637207, + "learning_rate": 2e-05, + "loss": 0.2129, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.8692678213119507, + "learning_rate": 2e-05, + "loss": 0.2059, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 5.119592666625977, + "learning_rate": 2e-05, + "loss": 0.3036, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 4.923192501068115, + "learning_rate": 2e-05, + "loss": 0.3127, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 6.533183574676514, + "learning_rate": 2e-05, + "loss": 0.4913, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 6.471742153167725, + "learning_rate": 2e-05, + "loss": 0.9363, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.208908557891846, + "learning_rate": 2e-05, + "loss": 0.8096, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.36978262662887573, + "learning_rate": 2e-05, + "loss": 0.072, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.794503211975098, + "learning_rate": 2e-05, + "loss": 0.2984, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.883040189743042, + "learning_rate": 2e-05, + "loss": 0.0473, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.8446871042251587, + "learning_rate": 2e-05, + "loss": 0.5398, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.6693472862243652, + "learning_rate": 2e-05, + "loss": 0.3187, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.1661856174468994, + "learning_rate": 2e-05, + "loss": 0.7084, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.1378469467163086, + "learning_rate": 2e-05, + "loss": 0.183, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.41687142848968506, + "learning_rate": 2e-05, + "loss": 0.147, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.9217529296875, + "learning_rate": 2e-05, + "loss": 0.2411, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.4083926677703857, + "learning_rate": 2e-05, + "loss": 0.4297, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.0639665126800537, + "learning_rate": 2e-05, + "loss": 0.4064, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0585795640945435, + "learning_rate": 2e-05, + "loss": 0.0417, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.4144511222839355, + "learning_rate": 2e-05, + "loss": 0.5697, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.6747444868087769, + "learning_rate": 2e-05, + "loss": 0.1211, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.0870719701051712, + "learning_rate": 2e-05, + "loss": 0.0471, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.8676738142967224, + "learning_rate": 2e-05, + "loss": 0.0567, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.028902441263198853, + "learning_rate": 2e-05, + "loss": 0.0803, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.1633816957473755, + "learning_rate": 2e-05, + "loss": 0.0649, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2211830323740672.0, + "train_loss": 0.2594273853302002, + "train_runtime": 62.8726, + "train_samples_per_second": 6.362, + "train_steps_per_second": 1.591 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2211830323740672.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..36c1e948e938cbd99ad4fcee3a1cf3c07eacf163 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e48a87c3f20beee6b60537a6415c23ad32a5f152ea070f29158ca8b3dbfe379e +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..127d53730ba70fa67c22c5b80c69c225e2f15249 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e5f80671ef87c91f8d877283f4d64ae17f3922692498ab8bade1274e4c623bee +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..95612311eb90bab5636c2b3839d4a94b10c4fb92 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d3053f8ef72499ec5eb87b19dde3585fe36d683ef1303e3895b3c1b01a3e5e1f +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b04c1022396bd0497d85af9a6dcc40b11db90cef --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc5926d983c1e7e9a2dc640798635fc1feee494f9fc9b42dfb147bfef2dfc573 +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..d8de77db964f08c2795bcf95fe2145fd714d9507 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:00bb472e961f7ccd8c1e04400b16d4c82a06a858239f380d6b49b2fc2dbe5e34 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..a131608df35a9edb193883e1f7846969c9d4ee05 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:10f9499a482921daed3b37b83c686bd84a16f3c3f234bb19d78d954d9dfd7d7c +size 368444338 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..4b5f878cd4b8f6c5af1088bf4560856454d892be --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b1206d09d7eecfb402e09ba78456b122e7e4de3ad4cb21f264c23ff8afc56e40 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..1a324151cba48f6258d13863258a5ee371c6568f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:86f7599824d58e4da2ff8d02ddba5020c1c9668d9e41eee9035f69d9c8f057cc +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..e646e2d26b1c1fd2bcd35c68f42ba0ebb37162c7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/18_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.05466795340180397, + "learning_rate": 2e-05, + "loss": 0.0016, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.7063548564910889, + "learning_rate": 2e-05, + "loss": 0.0618, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.35168468952178955, + "learning_rate": 2e-05, + "loss": 0.0128, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.05616375058889389, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.08031021058559418, + "learning_rate": 2e-05, + "loss": 0.4052, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 5.48492431640625, + "learning_rate": 2e-05, + "loss": 0.1711, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.6255172491073608, + "learning_rate": 2e-05, + "loss": 0.0994, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.22820666432380676, + "learning_rate": 2e-05, + "loss": 0.0351, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.8763604164123535, + "learning_rate": 2e-05, + "loss": 0.045, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.177014112472534, + "learning_rate": 2e-05, + "loss": 0.1762, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.058095533400774, + "learning_rate": 2e-05, + "loss": 0.1149, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3600158095359802, + "learning_rate": 2e-05, + "loss": 0.1317, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.271526962518692, + "learning_rate": 2e-05, + "loss": 0.0723, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 4.347335338592529, + "learning_rate": 2e-05, + "loss": 0.4234, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.04349707067012787, + "learning_rate": 2e-05, + "loss": 0.0426, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 5.021970748901367, + "learning_rate": 2e-05, + "loss": 0.2058, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5921158790588379, + "learning_rate": 2e-05, + "loss": 0.0353, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.17028647661209106, + "learning_rate": 2e-05, + "loss": 0.2059, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.1064319610595703, + "learning_rate": 2e-05, + "loss": 0.0617, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.2225959748029709, + "learning_rate": 2e-05, + "loss": 0.0325, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.6350761651992798, + "learning_rate": 2e-05, + "loss": 0.3661, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.14630341529846191, + "learning_rate": 2e-05, + "loss": 0.0068, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.5837147235870361, + "learning_rate": 2e-05, + "loss": 0.4262, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.6369373798370361, + "learning_rate": 2e-05, + "loss": 0.1053, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.2324517965316772, + "learning_rate": 2e-05, + "loss": 0.1023, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.7291555404663086, + "learning_rate": 2e-05, + "loss": 0.1058, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.6085166931152344, + "learning_rate": 2e-05, + "loss": 0.203, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.18125221133232117, + "learning_rate": 2e-05, + "loss": 0.3036, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.18139062821865082, + "learning_rate": 2e-05, + "loss": 0.0723, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8019345998764038, + "learning_rate": 2e-05, + "loss": 0.0814, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.0057206423953175545, + "learning_rate": 2e-05, + "loss": 0.0265, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.7829449772834778, + "learning_rate": 2e-05, + "loss": 0.1022, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.02886315807700157, + "learning_rate": 2e-05, + "loss": 0.1248, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.9490000605583191, + "learning_rate": 2e-05, + "loss": 0.076, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 4.736812591552734, + "learning_rate": 2e-05, + "loss": 0.3836, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.08967191725969315, + "learning_rate": 2e-05, + "loss": 0.0094, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.5028371810913086, + "learning_rate": 2e-05, + "loss": 0.117, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.464789390563965, + "learning_rate": 2e-05, + "loss": 0.1508, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 4.168581485748291, + "learning_rate": 2e-05, + "loss": 0.249, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.2872008681297302, + "learning_rate": 2e-05, + "loss": 0.0113, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.73968768119812, + "learning_rate": 2e-05, + "loss": 0.2528, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.04312964901328087, + "learning_rate": 2e-05, + "loss": 0.1341, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.57293701171875, + "learning_rate": 2e-05, + "loss": 0.014, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.09055174142122269, + "learning_rate": 2e-05, + "loss": 0.0115, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.8689008951187134, + "learning_rate": 2e-05, + "loss": 0.159, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.5440471172332764, + "learning_rate": 2e-05, + "loss": 0.124, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.3143815994262695, + "learning_rate": 2e-05, + "loss": 0.3719, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 4.91234016418457, + "learning_rate": 2e-05, + "loss": 0.0917, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.01908990554511547, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.909580945968628, + "learning_rate": 2e-05, + "loss": 0.242, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2206568850391040.0, + "train_loss": 0.1352577829360962, + "train_runtime": 66.2427, + "train_samples_per_second": 6.038, + "train_steps_per_second": 1.51 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2206568850391040.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..9b81431824f75d7427ee22060b486b2eed035cad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c256dc81447690ab0ed760dc6703e9a5f972e821e6bb5174afccd9365112188 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..6359de17ec300cc8fc78bd393788f601f917ee39 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1b24c51c241012173aa8d6aecc77fbfcef3d4a86b7c49c67e8e3490c2fc71a6e +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..5688211d330feb34ef7b2291fa793896905d4780 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f80b07083dca940c3ce2c9d30a786ef7b6ba7b88b145d3e52aff18bac0beae74 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..efab5a7181ba399196ecaa8e5f608c31c1ef7e15 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b6618feee7bf9f37ccc8bd063c782546c21657489bdfc3110b13100969029166 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5b48163b2c42138bb383a5db1a3269b7949c03f1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e18a208d067f71e49cae224b1d2a6121b92d8b2a00bf624b6bc739d7bbd26576 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..8249a2ec6e6c177a27012d8f902c36924626aed7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b0200466cec85ee40aafcf86f117a1dc2289fc3848768f62d988d1a2c474bf22 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9f35dff458b192d0b67b1888f871c27007c1c1c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d1d1e36f1ed3c384310748917ff18e68aa5d69d5b5b7b0448512168559774a8 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ee0e3e76b50dd8b4b63818e92d33e7bfe2ea86a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:788b39bdcee90c2b1952bc4fc177346a1c8367d1aabf659d5d527cfbe66d321b +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c4d0d8ce4548551faf4437696def973b1322cb9a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/19_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.3041911721229553, + "learning_rate": 2e-05, + "loss": 0.0787, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.30435091257095337, + "learning_rate": 2e-05, + "loss": 0.0694, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.42556527256965637, + "learning_rate": 2e-05, + "loss": 0.0434, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.7292447090148926, + "learning_rate": 2e-05, + "loss": 0.1488, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.0927646160125732, + "learning_rate": 2e-05, + "loss": 0.3292, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.320820152759552, + "learning_rate": 2e-05, + "loss": 0.0448, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.8121518492698669, + "learning_rate": 2e-05, + "loss": 0.0683, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.45110759139060974, + "learning_rate": 2e-05, + "loss": 0.119, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.24342703819274902, + "learning_rate": 2e-05, + "loss": 0.2696, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.12088610231876373, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.8707820773124695, + "learning_rate": 2e-05, + "loss": 0.1066, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.8355741500854492, + "learning_rate": 2e-05, + "loss": 0.0514, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.4699316024780273, + "learning_rate": 2e-05, + "loss": 0.1338, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.2308714389801025, + "learning_rate": 2e-05, + "loss": 0.111, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.761664628982544, + "learning_rate": 2e-05, + "loss": 0.209, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.33252766728401184, + "learning_rate": 2e-05, + "loss": 0.4898, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8503375053405762, + "learning_rate": 2e-05, + "loss": 0.2272, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3420613408088684, + "learning_rate": 2e-05, + "loss": 0.1102, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.21938271820545197, + "learning_rate": 2e-05, + "loss": 0.0367, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.08440004289150238, + "learning_rate": 2e-05, + "loss": 0.0531, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.7978280186653137, + "learning_rate": 2e-05, + "loss": 0.0671, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.54156094789505, + "learning_rate": 2e-05, + "loss": 0.5411, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.45233577489852905, + "learning_rate": 2e-05, + "loss": 0.1361, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.27233076095581055, + "learning_rate": 2e-05, + "loss": 0.2757, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.20192326605319977, + "learning_rate": 2e-05, + "loss": 0.2433, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.6086478233337402, + "learning_rate": 2e-05, + "loss": 0.2196, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.26809221506118774, + "learning_rate": 2e-05, + "loss": 0.187, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.1012892946600914, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.5395312309265137, + "learning_rate": 2e-05, + "loss": 0.1605, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.9917735457420349, + "learning_rate": 2e-05, + "loss": 0.072, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.11056128144264221, + "learning_rate": 2e-05, + "loss": 0.1626, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.0073158740997314, + "learning_rate": 2e-05, + "loss": 0.2558, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.470820665359497, + "learning_rate": 2e-05, + "loss": 0.3212, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.04533271864056587, + "learning_rate": 2e-05, + "loss": 0.0714, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.6545588970184326, + "learning_rate": 2e-05, + "loss": 0.39, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.4193216562271118, + "learning_rate": 2e-05, + "loss": 0.1312, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.23041048645973206, + "learning_rate": 2e-05, + "loss": 0.0291, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.153687477111816, + "learning_rate": 2e-05, + "loss": 0.5201, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.075352430343628, + "learning_rate": 2e-05, + "loss": 0.1655, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.2238205671310425, + "learning_rate": 2e-05, + "loss": 0.1478, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.4245926141738892, + "learning_rate": 2e-05, + "loss": 0.3311, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.09046297520399094, + "learning_rate": 2e-05, + "loss": 0.0219, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0018208930268883705, + "learning_rate": 2e-05, + "loss": 0.0335, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.1064792349934578, + "learning_rate": 2e-05, + "loss": 0.2516, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.9194135665893555, + "learning_rate": 2e-05, + "loss": 0.0403, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.5658713579177856, + "learning_rate": 2e-05, + "loss": 0.1061, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.33841052651405334, + "learning_rate": 2e-05, + "loss": 0.0845, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.1735719442367554, + "learning_rate": 2e-05, + "loss": 0.1685, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.15782494843006134, + "learning_rate": 2e-05, + "loss": 0.0434, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.08362329006195068, + "learning_rate": 2e-05, + "loss": 0.2733, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293519879012352.0, + "train_loss": 0.16353089809417726, + "train_runtime": 95.7537, + "train_samples_per_second": 4.177, + "train_steps_per_second": 1.044 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293519879012352.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..4590433cdd423fb34c43a62af66dd29312e5e7ae --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8b50570b26bba61fc40d01338c36013453803ce5ef418edc72db3c2d91266f8b +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..0d9b69f3e2aab984a652336bb12403148bc63ebc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:607bc96db229db023e522a3b012c7f0d545a0134af2e1b6aaba058aea16f9627 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..87311ff4a13786164a716045e79afe29b9dfb4f4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:99356faa8f99844dd5e4630750087f40bff128ba75bd312fc28ddcd9253dfd7c +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..ea7ca0aae2f32cb419d6816101b3757f87965e74 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4a19ba30592fcb1ed87982585c4ef844f328a95fa35d1b29814ec9c5fb5027a2 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5857fa6e5967c3ddaee28f959d355d35d0dca059 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:89384d845830d09c2a447b7ebe48c8e9983a09d9245ac61a9b2ae4abc9b729b8 +size 368442474 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6a9b58ceaa408be04759f4e52ce65018e7b54f0f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d78088cb3dc232c7c565e8cc84c8b5a317cf131637e581a9b51621dba9e503d1 +size 368443438 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..c7a8c6e6490272c03ea33a53f2f623e20bdc364c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4a8c58e5328a9ed10029c709be9ffe5a203aff41bd7ae9aaba057f9f79ebbd3a +size 368442474 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..621eda75a19da5897d7d944a01816233424e6aa7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a085fccec602984ce37941a2baeb0fd1a3303539fc215cecd8840b721bd0ae3c +size 368442474 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..9fb1c871f0bcbe586d70297a05082994d7b2dc41 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/1_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.030042845755815506, + "learning_rate": 2e-05, + "loss": 0.0169, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.04493667557835579, + "learning_rate": 2e-05, + "loss": 0.0401, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.17430752515792847, + "learning_rate": 2e-05, + "loss": 0.005, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.10912972688674927, + "learning_rate": 2e-05, + "loss": 0.0083, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.08338020741939545, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.6959218382835388, + "learning_rate": 2e-05, + "loss": 0.0374, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.007870294153690338, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.09560559689998627, + "learning_rate": 2e-05, + "loss": 0.0403, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.05684217810630798, + "learning_rate": 2e-05, + "loss": 0.0718, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.0033909156918525696, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.01980670727789402, + "learning_rate": 2e-05, + "loss": 0.0326, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.09405945986509323, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.68649423122406, + "learning_rate": 2e-05, + "loss": 0.088, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.024327317252755165, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.04908490553498268, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.16189368069171906, + "learning_rate": 2e-05, + "loss": 0.0354, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.9491567611694336, + "learning_rate": 2e-05, + "loss": 0.0255, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.0015416329260915518, + "learning_rate": 2e-05, + "loss": 0.0001, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.04257403314113617, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.7657027840614319, + "learning_rate": 2e-05, + "loss": 0.017, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.11868175119161606, + "learning_rate": 2e-05, + "loss": 0.0025, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.3353335857391357, + "learning_rate": 2e-05, + "loss": 0.1681, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.012815206311643124, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.022613856941461563, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.005410331301391125, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.017490020021796227, + "learning_rate": 2e-05, + "loss": 0.0213, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.7256666421890259, + "learning_rate": 2e-05, + "loss": 0.0492, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.01514490693807602, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.006880992092192173, + "learning_rate": 2e-05, + "loss": 0.074, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.007873348891735077, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.06169436499476433, + "learning_rate": 2e-05, + "loss": 0.001, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.012332312762737274, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.03053584136068821, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.14013749361038208, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.00361087778583169, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.3683127462863922, + "learning_rate": 2e-05, + "loss": 0.0065, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.7314853668212891, + "learning_rate": 2e-05, + "loss": 0.0161, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.007743148598819971, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.0979700088500977, + "learning_rate": 2e-05, + "loss": 0.0256, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.00447928998619318, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.006749676540493965, + "learning_rate": 2e-05, + "loss": 0.0013, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0007921429350972176, + "learning_rate": 2e-05, + "loss": 0.0002, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.023220648989081383, + "learning_rate": 2e-05, + "loss": 0.0881, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.011149675585329533, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.13010942935943604, + "learning_rate": 2e-05, + "loss": 0.0033, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.3392384648323059, + "learning_rate": 2e-05, + "loss": 0.0034, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.006402903702110052, + "learning_rate": 2e-05, + "loss": 0.0005, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.37901780009269714, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.007819139398634434, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.006767177954316139, + "learning_rate": 2e-05, + "loss": 0.0002, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2217003838341120.0, + "train_loss": 0.018431691527366637, + "train_runtime": 62.0077, + "train_samples_per_second": 6.451, + "train_steps_per_second": 1.613 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2217003838341120.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..24ec6a4e045871b1757cfa89e2fd40f4f6e7eac4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9b433988c9f5adac02eff4d841b93ecad1afcf96ddc6d1b2c392014d4917f69b +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..edd8455ff159dfbc98c2ecf3d253c15529ddb1ae --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e01d46b70f6d9ed0e3416de0a3a2fb76b9cacb9fe12b1d40ef9c0943c3197a29 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..6c70ec4f1501689f2b9ada3efe4cd7aa99684688 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b955acfcab7b3e0db05cb8545e88aa162835836602aa582d5289d252c9cba8d +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..4c9881666d5848c152116272bc2af0cab7012a7e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ab041304b444621008958882e10e662c59afd668f73de4b4d441ca853d9ef221 +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..ddf5fcf8f571d82e59d27d9d363ec9d525afc815 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4a87f0fed02b32e3512d894ff8e62dc8629bad298a4690aabdde629f15ae6bba +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b43ff19cfeb7564a4f28b7891a8983bf0aee1934 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6305a297855351ff188529c031fad060d884271f136a95bb16de36e53a18e36f +size 791579754 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..12b4b5606bfaf3bf1bebcccdb88d89ff486cc081 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3392e953f3b68ca0acb16201c0a512240834ef21134f2291f98ff60f6eface76 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..eca2f243e5ef115390312d45fdb874644df6ffba --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8b5d6c0d213511edbb62b0bdd58940a975566ed9130551e0c1469cbfcdb88dc8 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..4ab157c258f4b7523b06c5766b132b7025883298 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/20_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.08173560351133347, + "learning_rate": 2e-05, + "loss": 0.0322, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.6406750679016113, + "learning_rate": 2e-05, + "loss": 0.0884, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.8644487857818604, + "learning_rate": 2e-05, + "loss": 0.0872, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.4677561819553375, + "learning_rate": 2e-05, + "loss": 0.0259, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.332083225250244, + "learning_rate": 2e-05, + "loss": 0.4374, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.06755223870277405, + "learning_rate": 2e-05, + "loss": 0.0058, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.29981714487075806, + "learning_rate": 2e-05, + "loss": 0.047, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.9292137622833252, + "learning_rate": 2e-05, + "loss": 0.0791, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.685615062713623, + "learning_rate": 2e-05, + "loss": 0.2084, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.12010213732719421, + "learning_rate": 2e-05, + "loss": 0.1759, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.26363250613212585, + "learning_rate": 2e-05, + "loss": 0.0126, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.22038832306861877, + "learning_rate": 2e-05, + "loss": 0.016, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.9364656209945679, + "learning_rate": 2e-05, + "loss": 0.2828, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.14823439717292786, + "learning_rate": 2e-05, + "loss": 0.0124, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.9221945405006409, + "learning_rate": 2e-05, + "loss": 0.0309, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.5006041526794434, + "learning_rate": 2e-05, + "loss": 0.2352, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.08687044680118561, + "learning_rate": 2e-05, + "loss": 0.0444, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3566882312297821, + "learning_rate": 2e-05, + "loss": 0.0354, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2881341278553009, + "learning_rate": 2e-05, + "loss": 0.0688, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.4176519513130188, + "learning_rate": 2e-05, + "loss": 0.0634, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.30941125750541687, + "learning_rate": 2e-05, + "loss": 0.1483, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.8352800011634827, + "learning_rate": 2e-05, + "loss": 0.0694, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.2324819564819336, + "learning_rate": 2e-05, + "loss": 0.1269, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.7565291523933411, + "learning_rate": 2e-05, + "loss": 0.0411, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.9465610384941101, + "learning_rate": 2e-05, + "loss": 0.1573, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.32057100534439087, + "learning_rate": 2e-05, + "loss": 0.0123, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.4910975396633148, + "learning_rate": 2e-05, + "loss": 0.1215, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.059624284505844116, + "learning_rate": 2e-05, + "loss": 0.037, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.011060044169425964, + "learning_rate": 2e-05, + "loss": 0.1872, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09971418231725693, + "learning_rate": 2e-05, + "loss": 0.06, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.2357897758483887, + "learning_rate": 2e-05, + "loss": 0.2636, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.248442649841309, + "learning_rate": 2e-05, + "loss": 0.5686, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 5.4902496337890625, + "learning_rate": 2e-05, + "loss": 0.8532, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.0711212158203125, + "learning_rate": 2e-05, + "loss": 0.2079, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.6572651863098145, + "learning_rate": 2e-05, + "loss": 0.0555, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.06882038712501526, + "learning_rate": 2e-05, + "loss": 0.0327, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.14606298506259918, + "learning_rate": 2e-05, + "loss": 0.0157, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.255982875823975, + "learning_rate": 2e-05, + "loss": 0.6455, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.45053401589393616, + "learning_rate": 2e-05, + "loss": 0.0363, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.017607761546969414, + "learning_rate": 2e-05, + "loss": 0.0209, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.12663453817367554, + "learning_rate": 2e-05, + "loss": 0.0786, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.7460745573043823, + "learning_rate": 2e-05, + "loss": 0.0955, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.25482821464538574, + "learning_rate": 2e-05, + "loss": 0.0245, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.9458808898925781, + "learning_rate": 2e-05, + "loss": 0.1784, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.590663492679596, + "learning_rate": 2e-05, + "loss": 0.1577, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.5479228496551514, + "learning_rate": 2e-05, + "loss": 0.3116, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.5738053321838379, + "learning_rate": 2e-05, + "loss": 0.0399, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.05697108805179596, + "learning_rate": 2e-05, + "loss": 0.0036, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.1444709300994873, + "learning_rate": 2e-05, + "loss": 0.1636, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.07121957838535309, + "learning_rate": 2e-05, + "loss": 0.0067, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293985824243712.0, + "train_loss": 0.1342069113254547, + "train_runtime": 95.5527, + "train_samples_per_second": 4.186, + "train_steps_per_second": 1.047 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293985824243712.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..e0b2da91a6dd91748e3ba9c3dea1e3a77f3e2642 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:69c6aeea3a43cb59bfc7c95d1de601cb905fdde56f179c599e5c36eceb40d411 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..818151519ab489e776fdf3c77c942e2ee117646d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5fb99f89380d8cd6d25a8c61e0ce688155caec77ecbf43f0368f07eab525f3f4 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..42eab23df257a9f06d82f9fefe650db57e74f14f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6f8e8e11a514e0ddb0ff6b24ce9a50ad4d52bf77cf7fff52ea439f8b7a0e6361 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e2acb735fb94c4a2aad688d584c8b5dfcd307e50 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6fcbd8f634cd44fcee5259c45c8892b15cf36ea5769ced2af751fc68f1ecbb4 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..bc30614d4e98631846f441096211a15e8f396af1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97978cd6fe15795ed43ca0d905e95006af59924bfec2f26bd2f7e60347af25dc +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..2986a519e1769c10c903b2b179cd91b56a1de7ca --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ec93876d113bc39ab2118e868ea8bb3b2c841dfbf85f76267b2efb8b2a09c5e7 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..dc717f469ec2ab1300fbca3587f0293b15374552 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5578a935721a58ef17d1dd2fcdaf52363fd325986f79cec84805bba9036a4e5c +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..ce52128efeca041d284f615cdc50cb577079f6a2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a1d871f13b0270f0c5aa2aba45c0d680240dc9624753b53794f9fb2e9e3140b6 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..be5b8549a5602d2aabafa48ba0a499fc1b53bbe4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/2_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.3463573455810547, + "learning_rate": 2e-05, + "loss": 0.145, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.847277045249939, + "learning_rate": 2e-05, + "loss": 0.6353, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.1091071367263794, + "learning_rate": 2e-05, + "loss": 0.3744, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.0886058807373047, + "learning_rate": 2e-05, + "loss": 0.3593, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.0099315643310547, + "learning_rate": 2e-05, + "loss": 0.1566, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.711092472076416, + "learning_rate": 2e-05, + "loss": 0.481, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.39852067828178406, + "learning_rate": 2e-05, + "loss": 0.1602, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.812993049621582, + "learning_rate": 2e-05, + "loss": 0.2789, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.0793428421020508, + "learning_rate": 2e-05, + "loss": 0.2236, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.7409677505493164, + "learning_rate": 2e-05, + "loss": 0.4423, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.3502134680747986, + "learning_rate": 2e-05, + "loss": 0.0691, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.7947759628295898, + "learning_rate": 2e-05, + "loss": 0.2462, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.5871026515960693, + "learning_rate": 2e-05, + "loss": 0.5405, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.1611213684082031, + "learning_rate": 2e-05, + "loss": 0.1508, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.29902222752571106, + "learning_rate": 2e-05, + "loss": 0.0782, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.159346580505371, + "learning_rate": 2e-05, + "loss": 0.1872, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.4231855869293213, + "learning_rate": 2e-05, + "loss": 0.3242, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.5500972867012024, + "learning_rate": 2e-05, + "loss": 0.0537, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.479979157447815, + "learning_rate": 2e-05, + "loss": 0.5193, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 4.28127908706665, + "learning_rate": 2e-05, + "loss": 0.4141, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.0007649660110474, + "learning_rate": 2e-05, + "loss": 0.3391, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.5046669840812683, + "learning_rate": 2e-05, + "loss": 0.0719, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.064322590827942, + "learning_rate": 2e-05, + "loss": 0.2131, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.4081829786300659, + "learning_rate": 2e-05, + "loss": 0.1741, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.0727317333221436, + "learning_rate": 2e-05, + "loss": 0.1196, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.9382083415985107, + "learning_rate": 2e-05, + "loss": 0.2636, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.594624400138855, + "learning_rate": 2e-05, + "loss": 0.2708, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.2568050622940063, + "learning_rate": 2e-05, + "loss": 0.2373, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.0500417947769165, + "learning_rate": 2e-05, + "loss": 0.0644, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.094192735850811, + "learning_rate": 2e-05, + "loss": 0.1363, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.9045429229736328, + "learning_rate": 2e-05, + "loss": 0.3365, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.368138313293457, + "learning_rate": 2e-05, + "loss": 0.1746, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.6162572503089905, + "learning_rate": 2e-05, + "loss": 0.0756, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.5021755695343018, + "learning_rate": 2e-05, + "loss": 0.2408, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.7068824768066406, + "learning_rate": 2e-05, + "loss": 0.399, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.4401048421859741, + "learning_rate": 2e-05, + "loss": 0.4457, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5397492051124573, + "learning_rate": 2e-05, + "loss": 0.0435, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.848298966884613, + "learning_rate": 2e-05, + "loss": 0.2554, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.5884047746658325, + "learning_rate": 2e-05, + "loss": 0.0527, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.7609329223632812, + "learning_rate": 2e-05, + "loss": 0.4297, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.54994535446167, + "learning_rate": 2e-05, + "loss": 1.5117, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.20874109864234924, + "learning_rate": 2e-05, + "loss": 0.0124, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.7639610767364502, + "learning_rate": 2e-05, + "loss": 0.2961, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.2018752098083496, + "learning_rate": 2e-05, + "loss": 0.3223, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.493188858032227, + "learning_rate": 2e-05, + "loss": 1.0586, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.19529452919960022, + "learning_rate": 2e-05, + "loss": 0.1118, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.8132537603378296, + "learning_rate": 2e-05, + "loss": 0.2455, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.44210153818130493, + "learning_rate": 2e-05, + "loss": 0.1001, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.998722553253174, + "learning_rate": 2e-05, + "loss": 0.6489, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.47505417466163635, + "learning_rate": 2e-05, + "loss": 0.0773, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5221258941693952.0, + "train_loss": 0.2913630294799805, + "train_runtime": 105.3585, + "train_samples_per_second": 3.797, + "train_steps_per_second": 0.949 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5221258941693952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..2aa297fffd42de73d3d7895f47cc5819f617b9f6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1b8d99decbb876a7bd4f44429e1cd28aa65861ceeb61193a8df327c39dce486c +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..cfb1f3c0e85b4e9533928a77001716895a77e670 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:26c3bbe7476cf8da4845d91c6e69616d495a2a9ed4f5555e4d9bcb91db8d2188 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..6d2174bed36785e26d526c04b243c35a4e60f743 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3a8f74350fcd520e6486ce9d8c9aa2be74670dc781e962118608776794731e5b +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..53539819cae3ff2b303ac8379b321ca99002c308 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:22ef11c228360f0c0689f9ddddc10eb660bfdbfb2c0acb3cb6027bd7261c696f +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..fa37109b547a05060f87d9233cfb76886c8a99b9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2862c8f9a3e394dbe2bd982b671d2ac5a3cc719b995264cf56bd24e1400599a1 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..f4005eadc206f16044ace697713ec94a2e9189b2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:04d1226111dbcbfc4b68c68d9aecfa05bf795467ef9e5f25bb0286d06a1b1faf +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..0ea0fcd947b00688693c65134adb7425be948397 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:35304cd802d8df51ff4a1fe9bd3b12b54e1e4d1fdb4efb4892ceb9b87cd4a23d +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..e894d10816beffcba57a6683b4cd9ca6085579a7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:228215707b82c1de6abb7e2acd10852166e5a242fc4fba808049ec18e235f8ef +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..728f326fb9f329eb397abd225b3d01e170fd4c4f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/3_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.2841694355010986, + "learning_rate": 2e-05, + "loss": 0.8696, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.6060565710067749, + "learning_rate": 2e-05, + "loss": 0.3657, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.237285852432251, + "learning_rate": 2e-05, + "loss": 0.2932, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.222477436065674, + "learning_rate": 2e-05, + "loss": 0.6187, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.660562753677368, + "learning_rate": 2e-05, + "loss": 0.4898, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.8498005867004395, + "learning_rate": 2e-05, + "loss": 0.7941, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.085723400115967, + "learning_rate": 2e-05, + "loss": 0.7192, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.7799770832061768, + "learning_rate": 2e-05, + "loss": 0.5285, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.183253526687622, + "learning_rate": 2e-05, + "loss": 0.1643, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.9159126281738281, + "learning_rate": 2e-05, + "loss": 0.3895, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.7627202272415161, + "learning_rate": 2e-05, + "loss": 0.2723, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.4664148092269897, + "learning_rate": 2e-05, + "loss": 0.5059, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.3657524883747101, + "learning_rate": 2e-05, + "loss": 0.2239, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.1370654106140137, + "learning_rate": 2e-05, + "loss": 0.2851, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.3003127574920654, + "learning_rate": 2e-05, + "loss": 0.4018, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.776794195175171, + "learning_rate": 2e-05, + "loss": 0.3824, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.3813883066177368, + "learning_rate": 2e-05, + "loss": 0.6782, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.16526198387146, + "learning_rate": 2e-05, + "loss": 0.6274, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.1976152658462524, + "learning_rate": 2e-05, + "loss": 0.2133, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.9566861391067505, + "learning_rate": 2e-05, + "loss": 0.5073, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.556692361831665, + "learning_rate": 2e-05, + "loss": 0.5674, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.633335828781128, + "learning_rate": 2e-05, + "loss": 0.6268, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.9504110813140869, + "learning_rate": 2e-05, + "loss": 0.2206, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.5971667170524597, + "learning_rate": 2e-05, + "loss": 0.0665, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.686103343963623, + "learning_rate": 2e-05, + "loss": 0.5537, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.2432949542999268, + "learning_rate": 2e-05, + "loss": 0.709, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.08094106614589691, + "learning_rate": 2e-05, + "loss": 0.2074, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.0264437198638916, + "learning_rate": 2e-05, + "loss": 0.2625, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.7502768635749817, + "learning_rate": 2e-05, + "loss": 0.2917, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 3.2086546421051025, + "learning_rate": 2e-05, + "loss": 0.5125, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.985793113708496, + "learning_rate": 2e-05, + "loss": 0.3818, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.6222964525222778, + "learning_rate": 2e-05, + "loss": 0.2964, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.638920307159424, + "learning_rate": 2e-05, + "loss": 0.5143, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.4499127864837646, + "learning_rate": 2e-05, + "loss": 0.801, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.3849313259124756, + "learning_rate": 2e-05, + "loss": 0.3804, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.1666686534881592, + "learning_rate": 2e-05, + "loss": 0.6396, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.6014795303344727, + "learning_rate": 2e-05, + "loss": 0.6853, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.0441668033599854, + "learning_rate": 2e-05, + "loss": 0.2373, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.801032543182373, + "learning_rate": 2e-05, + "loss": 0.3555, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.4514472484588623, + "learning_rate": 2e-05, + "loss": 0.1724, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.6366243958473206, + "learning_rate": 2e-05, + "loss": 0.1365, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.0030248165130615, + "learning_rate": 2e-05, + "loss": 0.2479, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.4766112565994263, + "learning_rate": 2e-05, + "loss": 0.6035, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0275588035583496, + "learning_rate": 2e-05, + "loss": 0.2822, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5664529800415039, + "learning_rate": 2e-05, + "loss": 0.4925, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.9620707035064697, + "learning_rate": 2e-05, + "loss": 0.3365, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.07525634765625, + "learning_rate": 2e-05, + "loss": 0.2827, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.173569679260254, + "learning_rate": 2e-05, + "loss": 0.2039, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.1018731594085693, + "learning_rate": 2e-05, + "loss": 0.3497, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.6181703209877014, + "learning_rate": 2e-05, + "loss": 0.2529, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5408561442062336.0, + "train_loss": 0.4200180244445801, + "train_runtime": 116.6713, + "train_samples_per_second": 3.428, + "train_steps_per_second": 0.857 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5408561442062336.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..672be2e1db65987f9b512b772ec893e39b7a5e8c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c8df93ab9c093a98afbdf0e34a047694ab8be7f3727f68b312e44773e96b8ea +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..28093f135778e6eac7185b0e56c8379baa4bbd6e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ae2e323b831959ab1cd8cc73438cb7510c7384afc0240281e048bf9410dd0873 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..4f6d3cd478888d669e3c7513e4553e30c18586c2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3526e19ace5e811fe219a111797f84ee7735ae1f08908e9b55912af18d1b645 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..2ea0c905aa4a2b6cf87beb87d0add52045df4f8f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fdfbc60f9cb109010a4db1e1ce9ff5e2451063b43f4ea42d103d372f3fc63608 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..e3a21a1b938495d16c6c1040c0e4168e868a9f91 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1aabb61bef4882fa1b2511eeebaec4b16f99a5a494725b35a2eb9a911c96d415 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..beb3611bde348acdfd2c2eb0c6cea5e096affe35 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7358674cdc033731af5a92311c16e835d25a6ad8769e3c9cf2e7f9ef169f0f56 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..4b9e4902393fd32bedfb1ee050826b6e237679da --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4344ded0f93b620ca1ea51ee2264ca905a3ce4cd7fea7cfe7078c3b6088ab4c2 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..5cdb51e9c74290af947de8ac79a6791527cbe6e0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:37a271ff94240e2cd8b951ffe9dee067a040740ee707c8c845ea896066ddf10f +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..8c3cb800a0f23440218f5579be111866cfda783b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/4_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.32732781767845154, + "learning_rate": 2e-05, + "loss": 0.1906, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.105112910270691, + "learning_rate": 2e-05, + "loss": 0.3996, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.5827090740203857, + "learning_rate": 2e-05, + "loss": 0.3899, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.039531946182251, + "learning_rate": 2e-05, + "loss": 0.3489, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.6890852451324463, + "learning_rate": 2e-05, + "loss": 0.2225, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.3979721069335938, + "learning_rate": 2e-05, + "loss": 0.1274, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.5605970621109009, + "learning_rate": 2e-05, + "loss": 0.3015, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.1019654273986816, + "learning_rate": 2e-05, + "loss": 0.6064, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.2612954378128052, + "learning_rate": 2e-05, + "loss": 0.1975, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.4295094013214111, + "learning_rate": 2e-05, + "loss": 0.1774, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.8833520412445068, + "learning_rate": 2e-05, + "loss": 1.0161, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.8737852573394775, + "learning_rate": 2e-05, + "loss": 0.5062, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.418146014213562, + "learning_rate": 2e-05, + "loss": 0.3282, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.2475523948669434, + "learning_rate": 2e-05, + "loss": 0.4612, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.0454881191253662, + "learning_rate": 2e-05, + "loss": 0.3868, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.3253785371780396, + "learning_rate": 2e-05, + "loss": 0.3425, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.6156508922576904, + "learning_rate": 2e-05, + "loss": 0.207, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.7359356880187988, + "learning_rate": 2e-05, + "loss": 0.4717, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.6831457614898682, + "learning_rate": 2e-05, + "loss": 0.2229, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.662574529647827, + "learning_rate": 2e-05, + "loss": 0.6812, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.7348042726516724, + "learning_rate": 2e-05, + "loss": 0.4377, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.8819087743759155, + "learning_rate": 2e-05, + "loss": 0.2786, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.4360536336898804, + "learning_rate": 2e-05, + "loss": 0.3469, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.9159610271453857, + "learning_rate": 2e-05, + "loss": 0.3588, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.306160807609558, + "learning_rate": 2e-05, + "loss": 0.2275, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.0468307733535767, + "learning_rate": 2e-05, + "loss": 0.7288, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.9204359650611877, + "learning_rate": 2e-05, + "loss": 0.2908, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4894583225250244, + "learning_rate": 2e-05, + "loss": 0.1388, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.3722083568572998, + "learning_rate": 2e-05, + "loss": 0.3657, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.42203950881958, + "learning_rate": 2e-05, + "loss": 0.4248, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7556211352348328, + "learning_rate": 2e-05, + "loss": 0.2284, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.0730016231536865, + "learning_rate": 2e-05, + "loss": 0.5173, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.557481288909912, + "learning_rate": 2e-05, + "loss": 0.3608, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.1076911687850952, + "learning_rate": 2e-05, + "loss": 0.356, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.9837249517440796, + "learning_rate": 2e-05, + "loss": 0.3861, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.3684978485107422, + "learning_rate": 2e-05, + "loss": 0.2451, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.44278302788734436, + "learning_rate": 2e-05, + "loss": 0.2975, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.5889928340911865, + "learning_rate": 2e-05, + "loss": 0.5802, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.6001615524291992, + "learning_rate": 2e-05, + "loss": 0.3543, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.6057395935058594, + "learning_rate": 2e-05, + "loss": 0.287, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.7157738208770752, + "learning_rate": 2e-05, + "loss": 0.2234, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.2655953168869019, + "learning_rate": 2e-05, + "loss": 0.3411, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.941108226776123, + "learning_rate": 2e-05, + "loss": 0.1564, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.4111609160900116, + "learning_rate": 2e-05, + "loss": 0.2151, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.8199692964553833, + "learning_rate": 2e-05, + "loss": 0.3713, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.7880524396896362, + "learning_rate": 2e-05, + "loss": 0.2802, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.6125802993774414, + "learning_rate": 2e-05, + "loss": 0.4091, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.7019773721694946, + "learning_rate": 2e-05, + "loss": 0.356, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.5379943251609802, + "learning_rate": 2e-05, + "loss": 0.0908, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.4971425533294678, + "learning_rate": 2e-05, + "loss": 0.5535, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 6047770452426752.0, + "train_loss": 0.3558718872070312, + "train_runtime": 98.3738, + "train_samples_per_second": 4.066, + "train_steps_per_second": 1.017 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 6047770452426752.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..ec7d049acc67778f2c1fd1b8dfa581b74d29cfb8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:38b76c56357eb58357f1df8dad4851e9158bad43104be020fbc55d2d80f7f46d +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..cb38e3a65df07c1c092a8ecc24e54bbccc6c1532 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:439b522735b2512301fb7364d80ce57a2775c4054436db6697f4e38e56f8d79d +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..12f57983f049d545bc9012f9e890384b0be3c26f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f28f04b34953e320049f10ca584c6edde2c7d815cdc882ed8033332ab7ddda65 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..9e310d315520ce84d6368f119207ec230d37f7b1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a941e1d53acf57a77a63a684f3b1adbd0cb7515467148f90b509304dd3373cb4 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..590d051c06b404b2a40047974ef9184ca6ad1dab --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:94b48382b2c5b95bfdfd1ca0edc1f09451801f633313befb4b3d0a41e978757f +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..1f607e34993b28ef45402f7304ec1a4770059e96 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2d0cf7afd713ecd85533a895cdbebc5c6425feb60c1b11c7381e3a15cecf6be +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..1d9a0f4951167edbab96979ca720c8ab24e0e3d0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5a42dfaf0f47d9cfc9f493894307191c92f954d1506c65a4bbb94c62efdbb1d7 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..4476e1ea6dce03d1be882a57e5f7107e2982a9f6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:62968cbf8e8af1946e8207e2ac2d528599f42f78a51d3d4488dfe67e80610d3a +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..ff780e09bd9521fba626b311d860b71296a8157f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/5_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.9976876974105835, + "learning_rate": 2e-05, + "loss": 0.0679, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.7141196727752686, + "learning_rate": 2e-05, + "loss": 0.3051, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.0796979665756226, + "learning_rate": 2e-05, + "loss": 0.0819, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.17878755927085876, + "learning_rate": 2e-05, + "loss": 0.1108, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.5731177926063538, + "learning_rate": 2e-05, + "loss": 0.1197, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.1310681700706482, + "learning_rate": 2e-05, + "loss": 0.0396, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.007151928264647722, + "learning_rate": 2e-05, + "loss": 0.0369, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.0635114386677742, + "learning_rate": 2e-05, + "loss": 0.0232, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.4835498332977295, + "learning_rate": 2e-05, + "loss": 0.0946, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.290406346321106, + "learning_rate": 2e-05, + "loss": 0.1144, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.356996536254883, + "learning_rate": 2e-05, + "loss": 0.1253, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.8591537475585938, + "learning_rate": 2e-05, + "loss": 0.167, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.4957904815673828, + "learning_rate": 2e-05, + "loss": 0.1698, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.65490984916687, + "learning_rate": 2e-05, + "loss": 0.5781, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.1007421612739563, + "learning_rate": 2e-05, + "loss": 0.0364, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.7286807298660278, + "learning_rate": 2e-05, + "loss": 0.0449, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.6470589637756348, + "learning_rate": 2e-05, + "loss": 0.1085, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.22129499912261963, + "learning_rate": 2e-05, + "loss": 0.0349, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.7497376799583435, + "learning_rate": 2e-05, + "loss": 0.034, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.6989805698394775, + "learning_rate": 2e-05, + "loss": 0.1786, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.2352410852909088, + "learning_rate": 2e-05, + "loss": 0.4418, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.0415878295898438, + "learning_rate": 2e-05, + "loss": 0.265, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.26864564418792725, + "learning_rate": 2e-05, + "loss": 0.1844, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.14060114324092865, + "learning_rate": 2e-05, + "loss": 0.9526, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.261068820953369, + "learning_rate": 2e-05, + "loss": 0.2024, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.3645663261413574, + "learning_rate": 2e-05, + "loss": 0.1198, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.14812825620174408, + "learning_rate": 2e-05, + "loss": 0.0391, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.6551986932754517, + "learning_rate": 2e-05, + "loss": 0.1017, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.23034392297267914, + "learning_rate": 2e-05, + "loss": 0.0176, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8581911325454712, + "learning_rate": 2e-05, + "loss": 0.0536, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.678065299987793, + "learning_rate": 2e-05, + "loss": 0.064, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.310563325881958, + "learning_rate": 2e-05, + "loss": 0.6348, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.4946277439594269, + "learning_rate": 2e-05, + "loss": 0.0674, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.20147427916526794, + "learning_rate": 2e-05, + "loss": 0.0304, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.1072560548782349, + "learning_rate": 2e-05, + "loss": 0.2234, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.48706409335136414, + "learning_rate": 2e-05, + "loss": 0.2515, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.09151296317577362, + "learning_rate": 2e-05, + "loss": 0.2088, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.10992942750453949, + "learning_rate": 2e-05, + "loss": 0.0079, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.11463306844234467, + "learning_rate": 2e-05, + "loss": 0.5515, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.07315238565206528, + "learning_rate": 2e-05, + "loss": 0.0815, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.33847135305404663, + "learning_rate": 2e-05, + "loss": 0.0897, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0951613187789917, + "learning_rate": 2e-05, + "loss": 0.068, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.2293439507484436, + "learning_rate": 2e-05, + "loss": 0.0478, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.1636652946472168, + "learning_rate": 2e-05, + "loss": 0.0565, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.7970942258834839, + "learning_rate": 2e-05, + "loss": 0.0256, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.8482367992401123, + "learning_rate": 2e-05, + "loss": 0.2265, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.4550173878669739, + "learning_rate": 2e-05, + "loss": 0.0248, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 3.413026809692383, + "learning_rate": 2e-05, + "loss": 0.2727, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.027658531442284584, + "learning_rate": 2e-05, + "loss": 0.1944, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 5.384875297546387, + "learning_rate": 2e-05, + "loss": 1.1811, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5289199494234112.0, + "train_loss": 0.18316133975982665, + "train_runtime": 96.5815, + "train_samples_per_second": 4.142, + "train_steps_per_second": 1.035 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5289199494234112.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..3d1c33a572fa8c2b5f2b923984a7350dba6f7ebf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4f7d9eedab970357db8ae25678a02131194499166f3f5e8426a15e9ca00f935c +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..f47fc5fbb5bc72b3da3d9a8b64bef6748d922f62 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f00f8110a8efba4aa896be2ca63db8166e825ebccb499ed8e4af56949b56b47c +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..96ccb9d8785e047d467534c5089fcabf566a1f69 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fcf4d6b790ca6a279eed4a49200c7a021479c34d9ae2488de623231e6d5f1918 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..08fb693926f5ce1aa9380b3f9f25dc388dea60be --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:131fb4f55bb8d418e0df14645a085b4af069d730829a23a7f58b6cd5964a23b8 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..58268937f7115d02371c2740de45ddf4ea2b400d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c2719603f40edbf25325e5119b4c18c183e694d476f139682c806cb236799e17 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..59bd3573bfddf3faac4d99508b4b144422160bed --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3adf5055ba7f814ed82ea8bc80a6035d954d5f3b280f425ddd2067d6b6549161 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..3443e17c8a27f280ffec2c1abc0d9cdfab159a08 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:50001d33a5b070c91327bd0fd567233b5d9e834fbf66a9750a29fad2c747cd0b +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..32f6284704e9f6366c23ec1c577b30e20c869fda --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:24e038de7861126e069e2276ffaab84cd34ea7ab6a27aefeb3f2a2355326387f +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..02b6103fa9857fe6ca2e22ab5f03170cfd8518f2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/6_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.353604555130005, + "learning_rate": 2e-05, + "loss": 0.3323, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.8394666910171509, + "learning_rate": 2e-05, + "loss": 0.3833, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.4344497919082642, + "learning_rate": 2e-05, + "loss": 0.4331, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.9146625399589539, + "learning_rate": 2e-05, + "loss": 0.4968, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.9167978763580322, + "learning_rate": 2e-05, + "loss": 0.5498, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.8108319044113159, + "learning_rate": 2e-05, + "loss": 0.0524, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.722163200378418, + "learning_rate": 2e-05, + "loss": 0.3058, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.1330223083496094, + "learning_rate": 2e-05, + "loss": 0.5955, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.759941577911377, + "learning_rate": 2e-05, + "loss": 0.7412, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.3242213726043701, + "learning_rate": 2e-05, + "loss": 0.3839, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.0069427490234375, + "learning_rate": 2e-05, + "loss": 0.3291, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.5114264488220215, + "learning_rate": 2e-05, + "loss": 0.3058, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.5533579587936401, + "learning_rate": 2e-05, + "loss": 0.3553, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.218761682510376, + "learning_rate": 2e-05, + "loss": 0.2876, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.185713291168213, + "learning_rate": 2e-05, + "loss": 0.8413, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.318222761154175, + "learning_rate": 2e-05, + "loss": 0.3805, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.325293779373169, + "learning_rate": 2e-05, + "loss": 0.418, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.085848093032837, + "learning_rate": 2e-05, + "loss": 0.5223, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.8640496730804443, + "learning_rate": 2e-05, + "loss": 0.3511, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.2380738258361816, + "learning_rate": 2e-05, + "loss": 0.2921, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.4313981533050537, + "learning_rate": 2e-05, + "loss": 0.9524, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.8373114466667175, + "learning_rate": 2e-05, + "loss": 0.2379, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.315417528152466, + "learning_rate": 2e-05, + "loss": 0.3495, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.7658264636993408, + "learning_rate": 2e-05, + "loss": 0.3527, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.6111912727355957, + "learning_rate": 2e-05, + "loss": 0.3607, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.9168804883956909, + "learning_rate": 2e-05, + "loss": 0.3792, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.7781089544296265, + "learning_rate": 2e-05, + "loss": 0.3905, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.904635429382324, + "learning_rate": 2e-05, + "loss": 0.8389, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.1164379119873047, + "learning_rate": 2e-05, + "loss": 0.5055, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.1570987701416016, + "learning_rate": 2e-05, + "loss": 0.5917, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.7599525451660156, + "learning_rate": 2e-05, + "loss": 0.3618, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.2230757474899292, + "learning_rate": 2e-05, + "loss": 0.46, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.881989061832428, + "learning_rate": 2e-05, + "loss": 0.5205, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.2926905155181885, + "learning_rate": 2e-05, + "loss": 0.4578, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.7624220848083496, + "learning_rate": 2e-05, + "loss": 0.6611, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.9901823401451111, + "learning_rate": 2e-05, + "loss": 0.3514, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 6.04911994934082, + "learning_rate": 2e-05, + "loss": 0.8013, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.1320621967315674, + "learning_rate": 2e-05, + "loss": 0.5388, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.341943383216858, + "learning_rate": 2e-05, + "loss": 0.2848, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.8647819757461548, + "learning_rate": 2e-05, + "loss": 0.2039, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.8063185214996338, + "learning_rate": 2e-05, + "loss": 0.4224, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.740632951259613, + "learning_rate": 2e-05, + "loss": 0.6036, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.4234893321990967, + "learning_rate": 2e-05, + "loss": 0.2878, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0008856058120728, + "learning_rate": 2e-05, + "loss": 0.4297, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3499984741210938, + "learning_rate": 2e-05, + "loss": 0.2877, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.884402871131897, + "learning_rate": 2e-05, + "loss": 0.9565, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.2450984716415405, + "learning_rate": 2e-05, + "loss": 0.3003, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.617079257965088, + "learning_rate": 2e-05, + "loss": 0.3026, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.9369268417358398, + "learning_rate": 2e-05, + "loss": 0.4707, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.166783332824707, + "learning_rate": 2e-05, + "loss": 0.9443, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 1.0467000203608064e+16, + "train_loss": 0.4592701721191406, + "train_runtime": 149.314, + "train_samples_per_second": 2.679, + "train_steps_per_second": 0.67 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 1.0467000203608064e+16, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..03c7a4ce11f1a697731e06d7074219eadfa1e67c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8097466950cca3a89fe8ec4aa6a259d27c0fe83a0cb3ac792ca8c6c35e37f8e9 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..d6f5078e39488b8be7c11dccb9f60b2e338fb4aa --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d03002c32f44268beafdb60505f695ee50575ce3616ca7ed16e804478b179c34 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..cf76845a81fdb587a05e37100e5a591f9eb6e7a6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a081832b11d0977ca9d49a0db0cd49da12ad8dc071aa62b17740a5207a27d319 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..f5b14bb8667d48193a8586788159dc6bc44452e1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb71cba577619dd76b7d42b355d486ed28b881e9091c5a2049170b594fbbbd11 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..54e1628e8c17b7418a2a35ae9c9ec0767615096c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:712673f965e50680249bda3bf67c3c80c4f9cbe5614e29f316cbadc1eef0e5e5 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..99777a21cfdb0d243692ee98868d9d1ae774ac0a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9660b3b657983823fb3c61e2df74a36a212fe0418a46a7b294d6d0c2ca6ad313 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..cef37d849cfd556491f985c494561ae36ff9967e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:82f167608aeb6703695250ef7137a765e7feaf95ef7590c3f10c64e59698e0b5 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..1fd7c30bd286dce2405040f49d105289e82f43ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29be3f390bc859b093366b040d8e3e69aee7d94a7f4c9611c51ad2f41f96e3c3 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..de709a5b77ef01302602dae013db0f8c0b0e9f4b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/7_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.0230932235717773, + "learning_rate": 2e-05, + "loss": 0.1284, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.936081886291504, + "learning_rate": 2e-05, + "loss": 0.4431, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.55954909324646, + "learning_rate": 2e-05, + "loss": 0.1656, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.6446375846862793, + "learning_rate": 2e-05, + "loss": 0.2158, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.8311657905578613, + "learning_rate": 2e-05, + "loss": 0.6149, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.08901761472225189, + "learning_rate": 2e-05, + "loss": 0.0442, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.3289755880832672, + "learning_rate": 2e-05, + "loss": 0.0848, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.3928859233856201, + "learning_rate": 2e-05, + "loss": 0.3903, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.720211386680603, + "learning_rate": 2e-05, + "loss": 0.1906, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.4535422325134277, + "learning_rate": 2e-05, + "loss": 0.3051, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.09205679595470428, + "learning_rate": 2e-05, + "loss": 0.2277, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.265754461288452, + "learning_rate": 2e-05, + "loss": 0.2862, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.445840448141098, + "learning_rate": 2e-05, + "loss": 0.0297, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.935485601425171, + "learning_rate": 2e-05, + "loss": 0.2905, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.40889787673950195, + "learning_rate": 2e-05, + "loss": 0.3091, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.5615555644035339, + "learning_rate": 2e-05, + "loss": 0.2689, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.4605041742324829, + "learning_rate": 2e-05, + "loss": 0.0348, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.1703319549560547, + "learning_rate": 2e-05, + "loss": 0.423, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.6384815573692322, + "learning_rate": 2e-05, + "loss": 0.0622, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.514798641204834, + "learning_rate": 2e-05, + "loss": 0.6449, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.4533229470252991, + "learning_rate": 2e-05, + "loss": 0.4875, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.1554434299468994, + "learning_rate": 2e-05, + "loss": 0.2318, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.273431420326233, + "learning_rate": 2e-05, + "loss": 0.3099, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.1955086886882782, + "learning_rate": 2e-05, + "loss": 0.0103, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.3965133726596832, + "learning_rate": 2e-05, + "loss": 0.4299, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.5706011652946472, + "learning_rate": 2e-05, + "loss": 0.4178, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.031144576147198677, + "learning_rate": 2e-05, + "loss": 0.2868, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5437702536582947, + "learning_rate": 2e-05, + "loss": 0.0381, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.14137491583824158, + "learning_rate": 2e-05, + "loss": 0.0954, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 6.791753768920898, + "learning_rate": 2e-05, + "loss": 1.1566, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.3718093931674957, + "learning_rate": 2e-05, + "loss": 0.0716, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.5938986539840698, + "learning_rate": 2e-05, + "loss": 0.5088, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.356748104095459, + "learning_rate": 2e-05, + "loss": 0.1325, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.4490825831890106, + "learning_rate": 2e-05, + "loss": 0.1296, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.9598157405853271, + "learning_rate": 2e-05, + "loss": 0.3985, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.4340357780456543, + "learning_rate": 2e-05, + "loss": 0.3092, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.7279062271118164, + "learning_rate": 2e-05, + "loss": 0.5061, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.31227928400039673, + "learning_rate": 2e-05, + "loss": 0.3777, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.3051092624664307, + "learning_rate": 2e-05, + "loss": 0.3764, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.0521605014801025, + "learning_rate": 2e-05, + "loss": 0.8063, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.5782283544540405, + "learning_rate": 2e-05, + "loss": 0.1793, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.1769461631774902, + "learning_rate": 2e-05, + "loss": 0.1234, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.51096773147583, + "learning_rate": 2e-05, + "loss": 0.3944, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.8803030252456665, + "learning_rate": 2e-05, + "loss": 0.3628, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.4004777669906616, + "learning_rate": 2e-05, + "loss": 0.2554, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.800997793674469, + "learning_rate": 2e-05, + "loss": 0.1673, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.7868992686271667, + "learning_rate": 2e-05, + "loss": 0.1038, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.2046438753604889, + "learning_rate": 2e-05, + "loss": 0.0694, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.9411832094192505, + "learning_rate": 2e-05, + "loss": 0.1724, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.8009448051452637, + "learning_rate": 2e-05, + "loss": 0.267, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5500391349288960.0, + "train_loss": 0.2867129135131836, + "train_runtime": 97.9376, + "train_samples_per_second": 4.084, + "train_steps_per_second": 1.021 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5500391349288960.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5cdc4ec7035507879494d8ce72d015648c7b3ed6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d24a002184010c7fa9ca6e2fceda06eac62c9d323e11d8b8ce1fd4358dd472aa +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..e64880858416f6bbba3407ec5f41f68df1b4628c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:88a25ca601fbc26c9d53d4a8210b35661dd9d2c0cfb86b5e5e9a2bd5eb7c53d7 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..1daa845598a48f51a9365ad97d4cc8d0b8a51aa6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c0844774ff212a34378d1fb5207af379918728872d1056ed7e54d1edab41e5c6 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d7eb5853fe9b81b0749b217917c260f35d46fd3d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ef5e886f17593f51dca72a9e7542fdd80b05266623b27b5014edcd08b57709f9 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..b9c747fff4b02f0d4bb99afb672447658bc10374 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0eecc71902f385deec11f1a81c65ba13b70df59ff9f5e93cc8a517f200833a53 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..434722e39b5320667cd1cfea9fdd41e7e425372c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7db9caa5358943c182aa4b06195ecf8c0d4b2c693acdb6a5301c85bc5f86067b +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..384b7cb6615a46d70670d5357bb93f3b000a4d75 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:75018684f3c7ae585ebaec72ab271255c14eb0c938ae69f755d35b59379a275b +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..8fca9b87ba74f54329e83747b338e18d5b9575d4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23862ba0997b3cf4ff99c71cc8b6b29786a6f6a2b4881b6ff82afef94b4813ca +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..93fd411f6608ac3c242ca0f82a6467e7a5b6a2b6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/8_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.048345573246479034, + "learning_rate": 2e-05, + "loss": 0.0268, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.058467380702495575, + "learning_rate": 2e-05, + "loss": 0.0747, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.3030301630496979, + "learning_rate": 2e-05, + "loss": 0.0439, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.476255178451538, + "learning_rate": 2e-05, + "loss": 0.2644, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.036805130541324615, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.507547378540039, + "learning_rate": 2e-05, + "loss": 0.0779, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.05829893797636032, + "learning_rate": 2e-05, + "loss": 0.1209, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 4.811889171600342, + "learning_rate": 2e-05, + "loss": 0.2626, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.0734199285507202, + "learning_rate": 2e-05, + "loss": 0.1102, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.501260995864868, + "learning_rate": 2e-05, + "loss": 2.322, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.1522269248962402, + "learning_rate": 2e-05, + "loss": 0.4864, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.8252282738685608, + "learning_rate": 2e-05, + "loss": 0.3451, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.5864123106002808, + "learning_rate": 2e-05, + "loss": 0.1052, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.7750410437583923, + "learning_rate": 2e-05, + "loss": 0.0536, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.7311358451843262, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.037973105907440186, + "learning_rate": 2e-05, + "loss": 0.0083, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.7762229442596436, + "learning_rate": 2e-05, + "loss": 0.6668, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.730559825897217, + "learning_rate": 2e-05, + "loss": 0.3318, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.4120547771453857, + "learning_rate": 2e-05, + "loss": 0.5552, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.9716057777404785, + "learning_rate": 2e-05, + "loss": 0.5189, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 5.530657768249512, + "learning_rate": 2e-05, + "loss": 0.5571, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.001757284626364708, + "learning_rate": 2e-05, + "loss": 0.0876, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 3.1776905059814453, + "learning_rate": 2e-05, + "loss": 0.3716, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.267859697341919, + "learning_rate": 2e-05, + "loss": 0.2259, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.2935503423213959, + "learning_rate": 2e-05, + "loss": 0.2365, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.022208288311958313, + "learning_rate": 2e-05, + "loss": 0.1911, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.30496180057525635, + "learning_rate": 2e-05, + "loss": 0.2577, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4631206691265106, + "learning_rate": 2e-05, + "loss": 0.2643, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.8812381029129028, + "learning_rate": 2e-05, + "loss": 0.0617, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.0705350786447525, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 1.9453612565994263, + "learning_rate": 2e-05, + "loss": 0.2144, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.405311346054077, + "learning_rate": 2e-05, + "loss": 0.562, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.536633014678955, + "learning_rate": 2e-05, + "loss": 0.0691, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.5264396071434021, + "learning_rate": 2e-05, + "loss": 0.0245, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.361662894487381, + "learning_rate": 2e-05, + "loss": 0.0196, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.0302274227142334, + "learning_rate": 2e-05, + "loss": 0.1701, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.8131603002548218, + "learning_rate": 2e-05, + "loss": 0.0389, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.40075260400772095, + "learning_rate": 2e-05, + "loss": 0.0342, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.18031644821166992, + "learning_rate": 2e-05, + "loss": 0.0945, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.2239503115415573, + "learning_rate": 2e-05, + "loss": 0.021, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.5567588806152344, + "learning_rate": 2e-05, + "loss": 0.4756, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.6216201186180115, + "learning_rate": 2e-05, + "loss": 0.0264, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.17155307531356812, + "learning_rate": 2e-05, + "loss": 0.0092, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.5043379068374634, + "learning_rate": 2e-05, + "loss": 0.0362, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3260318040847778, + "learning_rate": 2e-05, + "loss": 0.1754, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.319258213043213, + "learning_rate": 2e-05, + "loss": 0.5061, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.20745891332626343, + "learning_rate": 2e-05, + "loss": 0.0081, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.860203504562378, + "learning_rate": 2e-05, + "loss": 0.157, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.16903749108314514, + "learning_rate": 2e-05, + "loss": 0.0099, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.10739284008741379, + "learning_rate": 2e-05, + "loss": 0.0077, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5327278363901952.0, + "train_loss": 0.22666221976280213, + "train_runtime": 102.2755, + "train_samples_per_second": 3.911, + "train_steps_per_second": 0.978 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5327278363901952.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..0cacfa1a023db1ed74ba01087b9ffd2a2873b990 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:906b4aae52700856dab23520de72bffb40e01a2f76f850dd89e5a69f69af33d5 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..6838d4f33d972de97e3af20d1199a2fd58547eb0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dee6485f8fdb1207bca5e50a1d82e94df9f396c4891ca708bc8be0a78e7c4b8b +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..da2b7727f8ac8db340ee4bf080d6ec86069f5c04 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9848a3895dfce9cb9b97e05638aa2f3b214769f9a7fd17c623d069eccb5e36d7 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..43aa4dc86dfe33b3f2e5ecb9762564a951ff6f19 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:138693b239dfbf895c2a21dcffca9e9fa2b11f3934474f1404d59fbe753b280d +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..4ce9846708a07dbd145b46e9e525778b7280dc2f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1ea97961ac3d0e018c199cd7f293c165106d4cff254253ad60923ae48d7440d7 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..3269941f1575e2814fbe00bac0d97f1f11ac2655 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d122e01a8575c9a036dda9aa462417be7c4feed9521126fb6f11d705303e7b0 +size 791578182 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..caa3c512de74d35f6286ff0a67849345d7753f2d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b0a25729c0d92b14ece8f7bc012d081883a919c90cf17313d9d5ff9e066aa0d4 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f5f9d5322cd6408d5cd4addded384d043e92d167 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f42ffc51c8cb529c90000ff66644f9fb1f336a37b34dc61dd1b6f070ebf0f0e7 +size 791576546 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..e0d20c6d6940ecb9f2c94617ed9e4625017b7fe1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/9_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.21795888245105743, + "learning_rate": 2e-05, + "loss": 0.3474, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.45443058013916016, + "learning_rate": 2e-05, + "loss": 0.6228, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.6299967169761658, + "learning_rate": 2e-05, + "loss": 0.07, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.179523229598999, + "learning_rate": 2e-05, + "loss": 0.1163, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.11994130164384842, + "learning_rate": 2e-05, + "loss": 0.3306, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.5945776700973511, + "learning_rate": 2e-05, + "loss": 0.147, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.056581445038318634, + "learning_rate": 2e-05, + "loss": 0.371, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.456613540649414, + "learning_rate": 2e-05, + "loss": 0.5596, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.46694907546043396, + "learning_rate": 2e-05, + "loss": 0.1613, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.9133310317993164, + "learning_rate": 2e-05, + "loss": 0.2544, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.1383259296417236, + "learning_rate": 2e-05, + "loss": 0.8061, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.715911388397217, + "learning_rate": 2e-05, + "loss": 0.4449, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.2591896057128906, + "learning_rate": 2e-05, + "loss": 0.5149, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.8498492240905762, + "learning_rate": 2e-05, + "loss": 0.0737, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.18627896904945374, + "learning_rate": 2e-05, + "loss": 0.1968, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.1116102784872055, + "learning_rate": 2e-05, + "loss": 0.0066, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5655573606491089, + "learning_rate": 2e-05, + "loss": 0.5681, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.8349328637123108, + "learning_rate": 2e-05, + "loss": 0.1849, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2748163342475891, + "learning_rate": 2e-05, + "loss": 0.0642, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.5974009037017822, + "learning_rate": 2e-05, + "loss": 0.3599, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.0770726203918457, + "learning_rate": 2e-05, + "loss": 0.2775, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.344602108001709, + "learning_rate": 2e-05, + "loss": 0.6187, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.5476309657096863, + "learning_rate": 2e-05, + "loss": 0.0508, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.1550190448760986, + "learning_rate": 2e-05, + "loss": 0.4132, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.4797477722167969, + "learning_rate": 2e-05, + "loss": 0.0712, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.5680513381958008, + "learning_rate": 2e-05, + "loss": 0.095, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.5048145055770874, + "learning_rate": 2e-05, + "loss": 0.5077, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.16661542654037476, + "learning_rate": 2e-05, + "loss": 0.0144, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.155276298522949, + "learning_rate": 2e-05, + "loss": 0.3728, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.6012856960296631, + "learning_rate": 2e-05, + "loss": 0.0806, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.1560981273651123, + "learning_rate": 2e-05, + "loss": 0.1894, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.274985671043396, + "learning_rate": 2e-05, + "loss": 0.2229, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.0257582664489746, + "learning_rate": 2e-05, + "loss": 0.2214, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7590439319610596, + "learning_rate": 2e-05, + "loss": 0.0844, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.1018344908952713, + "learning_rate": 2e-05, + "loss": 0.0189, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 4.628808975219727, + "learning_rate": 2e-05, + "loss": 0.459, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5059635043144226, + "learning_rate": 2e-05, + "loss": 0.2054, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.2131266593933105, + "learning_rate": 2e-05, + "loss": 0.2875, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.3236900568008423, + "learning_rate": 2e-05, + "loss": 0.1669, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.3683114051818848, + "learning_rate": 2e-05, + "loss": 0.4095, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.0751774311065674, + "learning_rate": 2e-05, + "loss": 0.3353, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.9231553077697754, + "learning_rate": 2e-05, + "loss": 0.2328, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.173947334289551, + "learning_rate": 2e-05, + "loss": 0.2192, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.2999845743179321, + "learning_rate": 2e-05, + "loss": 0.3794, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.6691272854804993, + "learning_rate": 2e-05, + "loss": 0.2465, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.53202223777771, + "learning_rate": 2e-05, + "loss": 0.2706, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.8418755531311035, + "learning_rate": 2e-05, + "loss": 0.5496, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.3514742851257324, + "learning_rate": 2e-05, + "loss": 0.0544, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.6382503509521484, + "learning_rate": 2e-05, + "loss": 0.2392, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.1557752788066864, + "learning_rate": 2e-05, + "loss": 0.0581, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291656215527424.0, + "train_loss": 0.27104862213134767, + "train_runtime": 97.5286, + "train_samples_per_second": 4.101, + "train_steps_per_second": 1.025 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291656215527424.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f3ce3ed682e6d70529a5cbb04045e208fc1126e3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bc3a1e90683c63278568ad7c9230d777c6399e4cb675fad1a1b20d5cd4292e2d +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..38aa84207fac91280f319f5fec35b274cb7fc2ea --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7ca56d149088d35df7a1795b4fb65f8d0d361986b8a8c6a9e9bce485d614cf6f +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_0/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..34ccc5101a02481bfc8c92f65a2c2d8a3034d3ac --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:43c68157b33553b349a86391a723dd2339da27bd85341e5baf558b63731d1971 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4b6700f160e16ed46117e32034946e16fa18db5b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:24f26aed8c48c4afce75c8408c5268fc07a9bd883a3ef25f8774df33257c88c8 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_1/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2bf63eecbc2627e5bde61b2339daeb24638b6356 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:086d6e516f7ca52675dada5b8ba21d8692e6cc244e3ffe2c48ad128c9b56155c +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..dfc12fcb3425b101509c47786002afad4a1e8e9b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76778aa59cfd028d47c067ef830e1b44fc6b3f71ba29809a1d2ab16d722b8d45 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_10/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..005d90d6b15611008331e32a4dd75b7704a76a26 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a1cb76f8180b96979e879e1b5f4b85ff75f1847e408477d452ca3ea9d1385094 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..815683463ed4c677cdf34a7d83388d0bbd52edfe --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b21981571daa4ddce41953d4980a71e2dfefa182d924c70478c83330abf1e938 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_11/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..61402a440c23efcfe4274881da40a885c911b415 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b3d6dc1480388889f8ef54dd710f1c20a6591183d61700e3511178e640712cd +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e7241267a1c84c123a3009ae744ef019e53321c4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a341b00e911e9fad7c322c31aa596525844d6e9dca0075cab9effd6986f40191 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_12/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..edaf740c620ca3b956e702eea6991491914f2397 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6fb548482fc88d8dd5149a79b89d7358a62e593d3bf1ec0f5ea232e30e62a696 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..92032d102753c511fd0a92d4031601df4265fde5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:21da1c597fc580aa0d3cc4beadce39678f1b15cf733bb65da69b279598b4612d +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_13/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0e8fb32a35a37c2a4e80e5626d15124aa5b4ccac --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f621edc3f437ae55ad726ec31349178a88945f243181a9a9ec267f4761e2fb30 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7ec1d627f837b4de396f24a772dc9af72fd9b551 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba54cb7041fc2b95eea97eb43d835010397caea82533795c06045e0398bebb31 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_14/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0fd7e1bb1bdea84e3ea99f9b0f2924452883473d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:960bc3f753ac7c6d40d8b73f2a7913fb1a523f22f7855b381a4a98c10bf22fa2 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6b87a82b29d9696598e44441cd2761b20bc0e11c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e6e21f172ec6efa9ef2d5f3e0e967cb3552faab188938411d8918157fd8b8de +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_15/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2822b0e9ad6c06280c838f2d8a63c619948fc056 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5ca40d030b22342d5aad9fee1ccda25f3f5efa3d43ed659875054a4af1f67d06 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..14fb0e145fe6030ae9554443b1a0880ff7d83e12 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:22a53d4e8c0451c1f69b446d8fe01b6aa5d8266262a72a641b5283213c00c9ee +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_16/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2ca443b53cc86117a4648c858497a2ae5f75a97e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6df1cc0fbd093504f68bd97cb58be4444f096e193d2c3a548469db03b2e35fdf +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e174af9e96334093dca509c1e065ddf74bcfdcea --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2a450cb039aede40ee38d3720a3e9e2a6795e5f515fe07948f81e71e8ca93015 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_17/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..38eb935f687d02e9f149c2d1c342144411454eff --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:016ad4acc977086a5e425064a7599963218eea100660250957b663ccc865e964 +size 833842160 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b616443e4e28e9d009df0a8397b57def8c619c84 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:df1f11ca9e58c133805bbb2151a61afae8c23c992c4bd199a95236ee18ccc5b8 +size 894036716 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_18/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..608b15d35526b5361be77963ed68b1fa1ab78398 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:090d17be1122edcb35983814a28e16a1eebe9df4ba993d982b220812fb32659d +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0ee46b305352e797c0119dac015b24c8c04b89c3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0c6f429303dad78861f497a700e8677e6294674dde94051b881b87d80e7ca581 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_19/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..95dc06fc508056022bbb870e48fa387f6c490112 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5f3a2b517d0b6a0936dd8de18ce1aac9e2127501e0556cf09c1fb0d1a35c9f77 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a5f21793eb345c622b47bfefc7bc3c173b1b7577 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:32ecfdad4565e2ae849ba8c0c3c801bb56c1b63e6158c173e621e1df6ef41978 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_2/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2974bfb95d5e0db7bebe9e2de4067df63274f08a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c344ed792feeb35069a60fc309ca8697f71470a4e1e5a3b787f6bc02c59a03c1 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..28205017010013fef09cd395e5a965c066efa273 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1362aa6def00ca4cc27be0977740d30632e949b3919b8b4a74e6b76d1a4cc031 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_20/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..999a58bc14485bffd7aec7ba1b83a435bc6440ca --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09af1b34dd59f0c27118d388867d9e9a0820739b82b10c0bd2d38e29f182bae9 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b0237e3647e13de34e08a74ed48a22a928e4e0c9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:214f68760977ba0c2196e749aa3b4b95fabb24e9a90802577168f4f325e84495 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_3/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6b5f44f11321ab489e20eadd14db5c8d4e65306b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ef1ea0c5e8f9d9767ba2cea814be0edfd150395862501935715e7e0a91128e39 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e11fc3815a7947d960856fa6106635a803ec9721 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5ec68ddc96cb9c316f1c613dd0ce262dce7493fe4c1f30d7779ea3e5e731064d +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_4/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..95c90ae68d1e156bb2b69809818076c4d70a1a1a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3f5f4313ae2f70061623c1e774ec4c4e63ab4262f9b24f33dd1f27d36e33f99 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a2b7f07d5e4f4bb47fd6d2eb829fb3e1ab626dd7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:642341c94659a01965ebc083cd18c5c331b6c60f5f633113f2a1b3a076e25349 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_5/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..580254ec4c67f34e7a85e880675df8ff3ff79b8c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:124191151a986a97dc4b24c65f7cf5c6da3bf369ce7eceb153a588d8e1f8571e +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a161738f17d4c177d5679cef721b7bf4f62aa4ee --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ccd0adfe16a3d7d511abae267a3ad1d082d883e85473f85150808a4708407296 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_6/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..fc19c4659ae3354dc4932076441836e3fea360af --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:898d47643081e372fbf933af7d2dfba3fb108dca1c0f2883c5f92e36e961431e +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0e0edc1cb584ac5470550bb0f98a083f46b2c92e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3e58d84b5c9f055001d046040d26d19a738ae22603b96d17e139001753ab1007 +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_7/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a24665caa503bc9bccc888ac7081ef6556d9b20e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0090fab23c3078685ee95fcf8bc7c3a6d2114749edc05fd495a40ab56aa78f96 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2d83fe03244d6d6a3c76a370b1b125d4d54bf273 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c15e14c4c378678be6996ff001095944a7f5c299819235257d1581a9c123119d +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_8/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..27e68b1b3eadeaa210eff3b81042ae9e87ba01ff --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c441dd218d8738f0bedb40abfb287a7ab8959d9fb36b8226ee45c7e5b9e64872 +size 2039612656 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..97fc52ef74c14a2db86be40e33e91c689f8aeb4f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c25d25c8e866cbdf6969696564fdc660b1057b709f630b1d81023e6dfe7cf9ef +size 1128606596 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/client_9/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round10_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round10_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..8480f3290fe02dcc8ad3f411dce4bb771607f136 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round10_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2a7a8d47edc80b205a0230dd0d312c16f3cc81aa0b01bd0d21718b6314735bb +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round11_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round11_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..120369764554a4d6cac58078d8512bf359607829 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round11_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf1dde35905b1ad125edf778aff1f9c3576ae0715c8c9876f665a63efcf18cd +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round12_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round12_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..64b454d945aa2a842c509c5cebe5ce0706d83a0c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round12_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb23ed1efc2080025ddd7ca376ff771b3f058afa24e048f5cb3cf99201aa1e03 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round13_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round13_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ad8d07954c41361e73a0c9c7d3b4b46d894b188 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round13_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8caca9dd439d57dbc5054f65d33ebbd7ac694940b2b9dbe5e8cdf68881137a2 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round14_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round14_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f05da1a0577050cea961f470b13842ab5f1d3f9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round14_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafca5706ac84f575d7648ee4786543561523e58cea8c01d2214b03d300bb988 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round15_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round15_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..351a4d2a284717aec90435f50067e1bb7002e9e4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round15_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:341a5a225c4c17a621d4f60caf4c34834000ea74e6f3c8b4f5aaada72ca1f547 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round16_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round16_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..08cfb19e58b68e297c605b87165d4acd241054e5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round16_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb378347913e2dd7be50fa6f9e3bd9e481f3897c2c9d31c597fe9164efc50bee +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round17_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round17_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..aeee08397c6123d33648f6703464379a4b7f85fe --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round17_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70c677ae195453774b8815440fee856ddb773c322681ac422ebc6e8abd5a58d7 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round18_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round18_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9935621fdc6da9e4abad379d97303b5ebf4f8f5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round18_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:393c3d5a15b40b2ff049c254f5b09f25d226b0971827b4e7460ca58f9ff7b76f +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round19_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round19_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..2766cf1a4da1390df01ea4146b4a3a7a06e1434e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round19_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dc77db94ac5b4746eebf2c850a9425da9a425a71108790d6b7f5f2654b614a8 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round1_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round1_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..34bfc5121e29fc587c3f574c17d8006189d3021a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round1_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47a94b88b812f21101345aae20a53cc10e0b346660c2cef669c77b201d06d508 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round20_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round20_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..4be93ba48f85d89c8cbd0e7426b79a10f51bc786 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round20_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20b9a2e67fdcb4051cebf52cba2ed0d47facb4dc69170f35a3483e3601bf25a4 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round2_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round2_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..b42d717305b30cd186c1c408b8719020c0658a5d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round2_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fd7458877767da0a36b80344cd6ec7eae8fbc8e11d07f8bfce3b4d5f5f239a9 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round3_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round3_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..01842d1316f21c3326c5735c6257b499da7ca458 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round3_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e537f056156082569ae0819abd06beb04b46581d0c63843bd68926e0ac2bff0e +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round4_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round4_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..3af9fe115499486166c3ef6d0a7bd27661ac08b4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round4_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a9649d6997f84508e5f566393e91e42e03a31aca55878bacd7ffd2e134bbfc +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round5_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round5_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..abd1e557c23a8769870b0beedaf7df93836a101e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round5_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7520bace4f501771d879621facd1b44f8db1ad7385a653e1535c7642465e67d +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round6_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round6_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..ebc846b7a5c587ced387cca85c4065e921c8cbe3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round6_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1951a910d1b7587459db2110e8ac43cbc906a34267372eb21d2244a900eaec46 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round7_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round7_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..5eb1370d3fa3177cb4fb710882f4b7eb706817e8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round7_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03f40a0df2437812b9e6de4e5d4e30014acb71344efa8132f19e1fc0f5925c82 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round8_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round8_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..385c599585a0973f240a027d28e985ca72d6b721 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round8_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3ef6d8247bc37b017895c9e5bbaf32f2db500f94f66b4ebade5a4a5ba7a271 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round9_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round9_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..219adb8927937a4597cd82efbf7ebb4c8a7519ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_iter100_round5_hetero/round9_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:277fa5f0607d993e80153171e4760f999c4ee6a65452ea62694d0b9b03818f22 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..1d0cb8c5963c48e05fd9d2cbece5e1b639e79466 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d05fa7197aa42cc14746069c90d7fc2eea469aacf40b37f97f4b4b925560e388 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..151a45f0a71c1cc934cf36bb0ccfe366fadc8bf2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c0aebfa6b1019a2f39179ec51dd0d5ed46a08b4935687bd0dddfe26e197af28 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..5da1f5bc168894482260af1667db1a7bcfe380c1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:31b803f1cb89f327c51756a11a47a24b6fe413018b83fad104d3afb57279dcbf +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e351f5687b3e6f8d294b9140187e839b71585d34 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1b8d4f3361fc516b78cb3078b68096902d0e2cd9ce2fb1eea8b91ff005d533ba +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..3ae083ecc8d92e06bfc0a98f183b2ecfded8fd5d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7ec6967d21b53de6be4edf4e5c17d7e263bfb653a9023ad7c5daa60aff243e95 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..b8f48587c9a2878f3ddec21c8b72c564a2d1bbee --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7a7e66ffd249c641309b0c4dc97f01fee6a03e2cf0555e1a53bd964bf4451e88 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..ed66613f8e3ada02ab6c29f5247334f2fb08cb75 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:672d5e960c95b45a3c3247e8d57f71015e46c7910f901af561bc8f2102307a35 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..0aacc9744833ecc232e97bcaa202839cdb689432 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:731c94cd2bec52f7423c622ac26942c0552e28735cff2e116ceadb00ddc738bc +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..6296da4848137e349b5795193b7aa8485c9385e4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/0_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.975966453552246, + "learning_rate": 2e-05, + "loss": 0.1873, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.7856782078742981, + "learning_rate": 2e-05, + "loss": 0.1761, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.3539968430995941, + "learning_rate": 2e-05, + "loss": 0.0319, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.20570360124111176, + "learning_rate": 2e-05, + "loss": 0.043, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.785933494567871, + "learning_rate": 2e-05, + "loss": 0.2768, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.799285888671875, + "learning_rate": 2e-05, + "loss": 0.267, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.8893604874610901, + "learning_rate": 2e-05, + "loss": 0.0934, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.229334831237793, + "learning_rate": 2e-05, + "loss": 0.456, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.8619229793548584, + "learning_rate": 2e-05, + "loss": 0.917, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.2525041103363037, + "learning_rate": 2e-05, + "loss": 0.2904, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.10519374907016754, + "learning_rate": 2e-05, + "loss": 0.0116, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.813767433166504, + "learning_rate": 2e-05, + "loss": 0.3906, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.22219926118850708, + "learning_rate": 2e-05, + "loss": 0.2138, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.19903819262981415, + "learning_rate": 2e-05, + "loss": 0.239, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.3674802780151367, + "learning_rate": 2e-05, + "loss": 0.2742, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.10247103124856949, + "learning_rate": 2e-05, + "loss": 0.0227, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.3424831628799438, + "learning_rate": 2e-05, + "loss": 0.1757, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 4.2065019607543945, + "learning_rate": 2e-05, + "loss": 0.3422, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.6023783087730408, + "learning_rate": 2e-05, + "loss": 0.2403, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.9348887801170349, + "learning_rate": 2e-05, + "loss": 0.0981, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.18112808465957642, + "learning_rate": 2e-05, + "loss": 0.0194, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.9743834733963013, + "learning_rate": 2e-05, + "loss": 0.2999, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4116685688495636, + "learning_rate": 2e-05, + "loss": 0.0843, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.1827506572008133, + "learning_rate": 2e-05, + "loss": 0.2052, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.29598069190979, + "learning_rate": 2e-05, + "loss": 0.0876, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.1951442956924438, + "learning_rate": 2e-05, + "loss": 0.0757, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.3302825391292572, + "learning_rate": 2e-05, + "loss": 0.2627, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.6185908317565918, + "learning_rate": 2e-05, + "loss": 0.1291, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.8732032775878906, + "learning_rate": 2e-05, + "loss": 0.5721, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09405488520860672, + "learning_rate": 2e-05, + "loss": 0.0061, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.6790041327476501, + "learning_rate": 2e-05, + "loss": 0.1661, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.5503239631652832, + "learning_rate": 2e-05, + "loss": 0.1043, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.30471983551979065, + "learning_rate": 2e-05, + "loss": 0.0166, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.9103224873542786, + "learning_rate": 2e-05, + "loss": 0.0828, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.635960340499878, + "learning_rate": 2e-05, + "loss": 0.2046, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.9675487279891968, + "learning_rate": 2e-05, + "loss": 0.0905, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.409539222717285, + "learning_rate": 2e-05, + "loss": 0.8484, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.6385936737060547, + "learning_rate": 2e-05, + "loss": 0.3508, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.5754489898681641, + "learning_rate": 2e-05, + "loss": 0.265, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.4736492931842804, + "learning_rate": 2e-05, + "loss": 0.3036, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.46990716457366943, + "learning_rate": 2e-05, + "loss": 0.0818, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.2309656143188477, + "learning_rate": 2e-05, + "loss": 0.1994, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.45787131786346436, + "learning_rate": 2e-05, + "loss": 0.0311, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0085583925247192, + "learning_rate": 2e-05, + "loss": 0.1179, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.117985725402832, + "learning_rate": 2e-05, + "loss": 0.1303, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.0631378889083862, + "learning_rate": 2e-05, + "loss": 0.2176, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.579545497894287, + "learning_rate": 2e-05, + "loss": 1.0527, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.0591726303100586, + "learning_rate": 2e-05, + "loss": 0.3907, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.421278715133667, + "learning_rate": 2e-05, + "loss": 0.28, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.22669081389904022, + "learning_rate": 2e-05, + "loss": 0.0608, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291730211438592.0, + "train_loss": 0.22968119621276856, + "train_runtime": 130.6318, + "train_samples_per_second": 3.062, + "train_steps_per_second": 0.766 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291730211438592.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..a0a4fd23ea6811630220d7815bba928b01e10aee --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2031506d84cac34526d7ed461e34e729a1a0ff824d1c9f3c70b27ece386eb39a +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..16e78adae0ee4689e6fcd0f2f602755b3f3aa675 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c9438d74fc073ef016f594e90c3405b5529e8439b804a66652403a08ff7a9209 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..9a31272b54a5b9987c306ba24c6b827ec4ed9414 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c1e5543921d145fa35180add1408e3d2580fa0bcd8c0b37798ea1af244b6fdc1 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..e48fbd8ace76b031d03d87e0b7f39721dc6ba244 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:65c42a9883c1761e767efeb090639bc08af9efce51e2274c52068e05cdc5b999 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..b1fd8a5b394189159f61ef80f14af8e20bdc49d6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:64acb2ccaac23d5e436401aed2a4fa28c85618fabc86755b1d96810bc55146fb +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..25b44f9334977e80a679bf78edbd99aef91de406 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5cc81e03a87ad9bb9ae395adfe78daa9eaa990cf66c789b6d256a813bfa6dc64 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..55348675344d6c5c04d8c93588148b87d1d6fc6a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2879733849a46dd64b3a39d74db67f60852ff92a48252bcf39b78258be3cfa5c +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..4d2e8e9756694de64058dbddc3e712c95ce46f3b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d27ab2edcacef56e6592d361204cf6f6af49b8ba7bb37c328a1e4e45f24c8db7 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c154f94795010ebcdfb43e24c13a0b77449cb252 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/10_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 4.277065277099609, + "learning_rate": 2e-05, + "loss": 0.37, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.201027870178223, + "learning_rate": 2e-05, + "loss": 0.17, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.6805689334869385, + "learning_rate": 2e-05, + "loss": 0.1426, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 3.4080328941345215, + "learning_rate": 2e-05, + "loss": 0.554, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.8445143699645996, + "learning_rate": 2e-05, + "loss": 0.2874, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.26823604106903076, + "learning_rate": 2e-05, + "loss": 0.0181, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 5.914226531982422, + "learning_rate": 2e-05, + "loss": 0.4081, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.797379732131958, + "learning_rate": 2e-05, + "loss": 0.1578, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.990204334259033, + "learning_rate": 2e-05, + "loss": 0.1063, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.9208881855010986, + "learning_rate": 2e-05, + "loss": 0.3017, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 3.6627352237701416, + "learning_rate": 2e-05, + "loss": 0.2138, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.6598092317581177, + "learning_rate": 2e-05, + "loss": 0.1494, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 4.791391849517822, + "learning_rate": 2e-05, + "loss": 0.5917, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 10.331496238708496, + "learning_rate": 2e-05, + "loss": 0.6699, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 9.474444389343262, + "learning_rate": 2e-05, + "loss": 0.5871, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 4.687495708465576, + "learning_rate": 2e-05, + "loss": 0.0981, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 4.005570411682129, + "learning_rate": 2e-05, + "loss": 0.1493, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.21403686702251434, + "learning_rate": 2e-05, + "loss": 0.2513, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 4.26255989074707, + "learning_rate": 2e-05, + "loss": 0.1963, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.352276563644409, + "learning_rate": 2e-05, + "loss": 0.3759, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.4846089482307434, + "learning_rate": 2e-05, + "loss": 0.0732, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.22196832299232483, + "learning_rate": 2e-05, + "loss": 0.6832, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.8316019773483276, + "learning_rate": 2e-05, + "loss": 0.0576, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.7157817482948303, + "learning_rate": 2e-05, + "loss": 0.2761, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 2.3334667682647705, + "learning_rate": 2e-05, + "loss": 0.1158, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.9155046939849854, + "learning_rate": 2e-05, + "loss": 0.7744, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 9.056317329406738, + "learning_rate": 2e-05, + "loss": 0.9888, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5096662640571594, + "learning_rate": 2e-05, + "loss": 0.0259, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 6.218349456787109, + "learning_rate": 2e-05, + "loss": 0.1187, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.913784980773926, + "learning_rate": 2e-05, + "loss": 0.4296, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 5.3420915603637695, + "learning_rate": 2e-05, + "loss": 0.4417, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.961092472076416, + "learning_rate": 2e-05, + "loss": 0.1435, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.622927665710449, + "learning_rate": 2e-05, + "loss": 0.8506, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.362285852432251, + "learning_rate": 2e-05, + "loss": 0.3533, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.482738733291626, + "learning_rate": 2e-05, + "loss": 0.0605, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 14.24290657043457, + "learning_rate": 2e-05, + "loss": 0.8999, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.250986576080322, + "learning_rate": 2e-05, + "loss": 0.2711, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 8.005436897277832, + "learning_rate": 2e-05, + "loss": 0.3023, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 4.815435886383057, + "learning_rate": 2e-05, + "loss": 0.7461, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.6406803131103516, + "learning_rate": 2e-05, + "loss": 0.7277, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 6.329057693481445, + "learning_rate": 2e-05, + "loss": 0.3978, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.07558009028434753, + "learning_rate": 2e-05, + "loss": 0.0059, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 6.101346969604492, + "learning_rate": 2e-05, + "loss": 0.3721, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 4.71949577331543, + "learning_rate": 2e-05, + "loss": 0.5647, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.1856529712677002, + "learning_rate": 2e-05, + "loss": 0.1434, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.344348907470703, + "learning_rate": 2e-05, + "loss": 0.4374, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 4.130119800567627, + "learning_rate": 2e-05, + "loss": 0.6937, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.6867356300354004, + "learning_rate": 2e-05, + "loss": 0.1486, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 4.034569263458252, + "learning_rate": 2e-05, + "loss": 0.3338, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.5193915963172913, + "learning_rate": 2e-05, + "loss": 0.4679, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2221568423886848.0, + "train_loss": 0.35407758712768556, + "train_runtime": 96.1523, + "train_samples_per_second": 4.16, + "train_steps_per_second": 1.04 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2221568423886848.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..31a585ca8f1f5a79a9b37446919bec06f8f0deec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:07d7e6abe8c7593c5b135ce64c73cfbbfff947a3d7ec4eb0046ba8155b97266d +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..b6e2148cbda7659c30dc414d4a515ba8eadf5120 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72fd5020decf86475e0a6b2e72a29db37e3f46cd73a95e00566da829531eadb3 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a9201355f44b65dd6038096d382f72b42a5b3dc3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:837768cf5239125c38d69b86ef3f4e2b722ade7af63716eb9d5c9a924a8e9c61 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..6a1faf9df7729ef64f9d55da7c6b16e9e1a78f2c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aa37e800198553e3312e3483ba89bbe778788fdf267864565681b7b68ccb2fd4 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..68931a03ee5ba38da4c2079b359fd6b2161a1b83 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a2333d48276706bd5b840ab606d2ee581c914cde5d5ea7be95ef3f6881ccbdb3 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..16fce2eb8cdaee6b8c228a049be82a20618cd33a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20afa5b9e4d4ccde8f62554f99a50a2f5da3518ed78d83d067eb1f71b81be0d8 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..a0fc8db515b846ba10a747df1afd534affd28589 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3a231d157e464a7a3159eb6be2dfde8f6eefeb67fc1f5573ec55052a8af60e7d +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f92b7ed453df6bc1ee3f01224946753991c8b171 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72b46af2f53ffdf5cc8bfd0f896933a13b3910d79963154632325ab34151e4d8 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..fe61d33f502cb4eb8beb9914a7bb74b28f17320d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/11_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.206027030944824, + "learning_rate": 2e-05, + "loss": 0.4937, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.521653175354004, + "learning_rate": 2e-05, + "loss": 0.4807, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.5993270874023438, + "learning_rate": 2e-05, + "loss": 0.3508, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.7889204025268555, + "learning_rate": 2e-05, + "loss": 0.3419, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.209856033325195, + "learning_rate": 2e-05, + "loss": 0.707, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.1698633432388306, + "learning_rate": 2e-05, + "loss": 0.3842, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.2927684783935547, + "learning_rate": 2e-05, + "loss": 0.8955, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 4.133338928222656, + "learning_rate": 2e-05, + "loss": 0.7778, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.1304954290390015, + "learning_rate": 2e-05, + "loss": 0.5952, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.9711688160896301, + "learning_rate": 2e-05, + "loss": 0.3536, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.0043792724609375, + "learning_rate": 2e-05, + "loss": 0.2974, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.6223782300949097, + "learning_rate": 2e-05, + "loss": 0.5062, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.3750557899475098, + "learning_rate": 2e-05, + "loss": 0.4, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 3.1023905277252197, + "learning_rate": 2e-05, + "loss": 0.3938, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.0721293687820435, + "learning_rate": 2e-05, + "loss": 0.4673, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.0541162490844727, + "learning_rate": 2e-05, + "loss": 0.418, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.0855062007904053, + "learning_rate": 2e-05, + "loss": 0.5259, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.2667620182037354, + "learning_rate": 2e-05, + "loss": 0.4634, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.8005095720291138, + "learning_rate": 2e-05, + "loss": 0.5481, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.6507138609886169, + "learning_rate": 2e-05, + "loss": 0.4591, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 3.3960788249969482, + "learning_rate": 2e-05, + "loss": 0.355, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.0557074546813965, + "learning_rate": 2e-05, + "loss": 0.5819, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.1364530324935913, + "learning_rate": 2e-05, + "loss": 0.4724, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.482191801071167, + "learning_rate": 2e-05, + "loss": 0.272, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 3.825197458267212, + "learning_rate": 2e-05, + "loss": 0.5956, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.0834991931915283, + "learning_rate": 2e-05, + "loss": 0.4834, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.6666842699050903, + "learning_rate": 2e-05, + "loss": 0.3306, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 2.192619562149048, + "learning_rate": 2e-05, + "loss": 0.2664, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.5968096256256104, + "learning_rate": 2e-05, + "loss": 0.2624, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.6724826097488403, + "learning_rate": 2e-05, + "loss": 0.3294, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.9662816524505615, + "learning_rate": 2e-05, + "loss": 0.4375, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.3559138774871826, + "learning_rate": 2e-05, + "loss": 0.3655, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 3.794806718826294, + "learning_rate": 2e-05, + "loss": 0.6509, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.1000509187579155, + "learning_rate": 2e-05, + "loss": 0.1748, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.9079816341400146, + "learning_rate": 2e-05, + "loss": 0.323, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.2479264736175537, + "learning_rate": 2e-05, + "loss": 0.4466, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.032731771469116, + "learning_rate": 2e-05, + "loss": 0.6792, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.7883899807929993, + "learning_rate": 2e-05, + "loss": 0.3509, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 4.1185808181762695, + "learning_rate": 2e-05, + "loss": 0.6039, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.401890277862549, + "learning_rate": 2e-05, + "loss": 0.6777, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.437748670578003, + "learning_rate": 2e-05, + "loss": 0.7432, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.8864333629608154, + "learning_rate": 2e-05, + "loss": 0.6943, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.663733720779419, + "learning_rate": 2e-05, + "loss": 0.3433, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.9872604608535767, + "learning_rate": 2e-05, + "loss": 0.4905, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.01814078539609909, + "learning_rate": 2e-05, + "loss": 0.4973, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.1219398975372314, + "learning_rate": 2e-05, + "loss": 0.3457, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.315551519393921, + "learning_rate": 2e-05, + "loss": 0.4995, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.306283473968506, + "learning_rate": 2e-05, + "loss": 0.4502, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.5388658046722412, + "learning_rate": 2e-05, + "loss": 0.3784, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.8539464473724365, + "learning_rate": 2e-05, + "loss": 0.4585, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2192326831112192.0, + "train_loss": 0.46838202953338626, + "train_runtime": 81.9962, + "train_samples_per_second": 4.878, + "train_steps_per_second": 1.22 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2192326831112192.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..b214c762fc4f321593c06c2eebd1a44246ca014d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:37a31a67a5662e84646c099bd1d0853b922173a1a9bc08ede28abc29696213a6 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..59bf63d097a8932cc1230753abb6ca4429d0462f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ab2a5739c9091396746022ff99aa4ccec43f7a2e63ec1158df4f3193028ae3ae +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..52fabea89a5a8159cd657806a17b206426d6ae40 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72050d0bdb2fd4af7dec3208f97802c874ab38442a718dac26a705fad6f67be6 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b5fa45c3429c38a26097fbef8ba7bdbf4a692e02 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4b174f44745eb06be0315457f56b21c3af9d6e92c6bc220361ffa5c31ec746a8 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..5db6463ecb0ba84695d4dfdad572f0be8f4f36f5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:75e8f59eda3dc4adebdf6782cdde5391191a3a1098db2046b6e00ed393d6b07a +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..438c079883d7fa9b4b969fa6bcf94fd699d7d614 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b04a70125c02d379b3dcadbd874267cd655f8bcd1c8d9b40b2c24a001402d85e +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2eed71654f8c91c2940e9563bcd12674959a9ff4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:de3cdcf996a10ac940c6d9ae557b26e825d867cc61fc19923b402db7391d308d +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..52130e8ec8585988d9de572a3571881ac6bd467a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8f7249b04c6f91af6fe98cb313ee73380106b605608a418085db738e50ae4ac +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..a1983c3ba253d8a356835283cbb5acc62552fe92 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/12_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.059566855430603, + "learning_rate": 2e-05, + "loss": 0.0448, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.0022856448777019978, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.07661868631839752, + "learning_rate": 2e-05, + "loss": 0.0143, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.5957555770874023, + "learning_rate": 2e-05, + "loss": 0.0277, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.5000383257865906, + "learning_rate": 2e-05, + "loss": 0.027, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.06301424652338028, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.159196376800537, + "learning_rate": 2e-05, + "loss": 0.5704, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.369729995727539, + "learning_rate": 2e-05, + "loss": 0.2958, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.2631317675113678, + "learning_rate": 2e-05, + "loss": 0.1629, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.24118709564209, + "learning_rate": 2e-05, + "loss": 0.1438, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.10603156685829163, + "learning_rate": 2e-05, + "loss": 0.1289, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.16504910588264465, + "learning_rate": 2e-05, + "loss": 0.0135, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.32078906893730164, + "learning_rate": 2e-05, + "loss": 0.0534, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.8100401163101196, + "learning_rate": 2e-05, + "loss": 0.1325, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.029811494052410126, + "learning_rate": 2e-05, + "loss": 0.0112, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.0602889358997345, + "learning_rate": 2e-05, + "loss": 0.0128, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.09126672148704529, + "learning_rate": 2e-05, + "loss": 0.0382, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.11353413015604019, + "learning_rate": 2e-05, + "loss": 0.0684, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.13520093262195587, + "learning_rate": 2e-05, + "loss": 0.0238, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.4951421916484833, + "learning_rate": 2e-05, + "loss": 0.0377, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.7304787635803223, + "learning_rate": 2e-05, + "loss": 0.2942, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.019591914489865303, + "learning_rate": 2e-05, + "loss": 0.036, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.432101845741272, + "learning_rate": 2e-05, + "loss": 0.199, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.5867564678192139, + "learning_rate": 2e-05, + "loss": 0.0275, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.05375500023365021, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.07455942779779434, + "learning_rate": 2e-05, + "loss": 0.0025, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.0313551239669323, + "learning_rate": 2e-05, + "loss": 0.0019, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.1275455355644226, + "learning_rate": 2e-05, + "loss": 0.0049, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.015496511943638325, + "learning_rate": 2e-05, + "loss": 0.0024, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.06546977162361145, + "learning_rate": 2e-05, + "loss": 0.0046, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.009432588703930378, + "learning_rate": 2e-05, + "loss": 0.0059, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.012232727371156216, + "learning_rate": 2e-05, + "loss": 0.0066, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.10806374996900558, + "learning_rate": 2e-05, + "loss": 0.0037, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.6549261212348938, + "learning_rate": 2e-05, + "loss": 0.0293, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.005222799722105265, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.04395957291126251, + "learning_rate": 2e-05, + "loss": 0.2404, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.0455283485352993, + "learning_rate": 2e-05, + "loss": 0.0026, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.008798840455710888, + "learning_rate": 2e-05, + "loss": 0.0187, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.3772972822189331, + "learning_rate": 2e-05, + "loss": 0.1992, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.0035173140931874514, + "learning_rate": 2e-05, + "loss": 0.0089, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.050068192183971405, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.07186979800462723, + "learning_rate": 2e-05, + "loss": 0.0052, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.0033018523827195168, + "learning_rate": 2e-05, + "loss": 0.0077, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.7036596536636353, + "learning_rate": 2e-05, + "loss": 0.1722, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.0036176389548927546, + "learning_rate": 2e-05, + "loss": 0.0146, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.24624799191951752, + "learning_rate": 2e-05, + "loss": 0.0595, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 3.340324640274048, + "learning_rate": 2e-05, + "loss": 0.3942, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.027079716324806213, + "learning_rate": 2e-05, + "loss": 0.4802, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.04889419674873352, + "learning_rate": 2e-05, + "loss": 0.0028, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.0563852787017822, + "learning_rate": 2e-05, + "loss": 0.0319, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5287535144075264.0, + "train_loss": 0.08146116137504578, + "train_runtime": 126.0562, + "train_samples_per_second": 3.173, + "train_steps_per_second": 0.793 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5287535144075264.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..941a9140ccb0fe8b174b9b2601db561b5e60ee28 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4471abf66c32bb5934211bdcaca977a0de18518196a599d0481bb43fd558cfb3 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..dc76a03faa4d2596e32ae6155d1b3481cfc61462 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:694f94f9c8ecfd58ca2fe041e1ed31e44f6f4ceb98743cbcdf1af5e70f3b143f +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..80e811d2b5c447c534e2289bf938b2b232644664 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:64b653bb6b84a231e268f55e1bcc41c68de991d8a2e8f77def3d7bbc96a420e2 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..696148e02b1afb0dbfad2cedf8b271ee2217c372 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eca53d07022846889cbf41f6921a2895eb3524f6d4c8136d74b6c2a6cdefd361 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..513d99c2d8c972cdf0d4836e8781571c7057d3d9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a0d94129e97289ab5af10d9ae76fc0f6ebbe886733f2e7e430256497136be912 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..9dc7d0b6934093a9fd7b3a161069072b6b695b4e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0edba7c9e75d8a7a38c00af91efc7767bcfa5754fa5b833654eeaad45d60d4b8 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..9fa333e37aae88f64897c028c05b95825beeeca8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0f5d587fa4bc7316566bafb0c7a88f0f3b2a7a3a9e6924d6fe054127a6bafb3c +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..38c85cda4059d1d84b8d219ed7880ea35b26579e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cf3ddcdb4ce2fbaab7c2379c452a0d96afddc82f9ece388971fdb1b3d0e31a15 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3bfab15db5e390ae689baf12a463f795055874bc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/13_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.801525115966797, + "learning_rate": 2e-05, + "loss": 0.2587, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.3922131061553955, + "learning_rate": 2e-05, + "loss": 0.2448, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.7673801183700562, + "learning_rate": 2e-05, + "loss": 0.4633, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.9585041403770447, + "learning_rate": 2e-05, + "loss": 0.087, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.6602948307991028, + "learning_rate": 2e-05, + "loss": 0.0905, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.17493529617786407, + "learning_rate": 2e-05, + "loss": 0.0312, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.7638754844665527, + "learning_rate": 2e-05, + "loss": 0.0955, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.501267671585083, + "learning_rate": 2e-05, + "loss": 0.4276, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.7278720140457153, + "learning_rate": 2e-05, + "loss": 0.2224, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.1080026626586914, + "learning_rate": 2e-05, + "loss": 1.252, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.27341890335083, + "learning_rate": 2e-05, + "loss": 0.498, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 5.594723701477051, + "learning_rate": 2e-05, + "loss": 0.2343, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.8931491374969482, + "learning_rate": 2e-05, + "loss": 0.9289, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.5089459419250488, + "learning_rate": 2e-05, + "loss": 0.1591, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.4253664016723633, + "learning_rate": 2e-05, + "loss": 0.2861, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.0828744173049927, + "learning_rate": 2e-05, + "loss": 0.1164, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.6505520939826965, + "learning_rate": 2e-05, + "loss": 0.2263, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.7339869737625122, + "learning_rate": 2e-05, + "loss": 0.3087, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.4004045724868774, + "learning_rate": 2e-05, + "loss": 0.1217, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.1830326318740845, + "learning_rate": 2e-05, + "loss": 0.2277, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.04140545800328255, + "learning_rate": 2e-05, + "loss": 0.2523, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.390353202819824, + "learning_rate": 2e-05, + "loss": 0.4085, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.025927145034074783, + "learning_rate": 2e-05, + "loss": 0.013, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.7939550876617432, + "learning_rate": 2e-05, + "loss": 0.1918, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.33892422914505005, + "learning_rate": 2e-05, + "loss": 0.247, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.0044050216674805, + "learning_rate": 2e-05, + "loss": 0.3637, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.5945000648498535, + "learning_rate": 2e-05, + "loss": 0.3102, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.8856892585754395, + "learning_rate": 2e-05, + "loss": 0.1621, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.07295922189950943, + "learning_rate": 2e-05, + "loss": 0.0444, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.09886308014392853, + "learning_rate": 2e-05, + "loss": 0.0757, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.9319190979003906, + "learning_rate": 2e-05, + "loss": 0.2114, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.3254298269748688, + "learning_rate": 2e-05, + "loss": 0.0883, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.1689889430999756, + "learning_rate": 2e-05, + "loss": 0.3984, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.501957654953003, + "learning_rate": 2e-05, + "loss": 0.2849, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.196662187576294, + "learning_rate": 2e-05, + "loss": 0.2261, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.32526344060897827, + "learning_rate": 2e-05, + "loss": 0.0443, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5529802441596985, + "learning_rate": 2e-05, + "loss": 0.5536, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.6129283308982849, + "learning_rate": 2e-05, + "loss": 0.0964, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.293200135231018, + "learning_rate": 2e-05, + "loss": 0.2651, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.0160406827926636, + "learning_rate": 2e-05, + "loss": 0.0863, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.6945836544036865, + "learning_rate": 2e-05, + "loss": 0.1134, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 3.7470624446868896, + "learning_rate": 2e-05, + "loss": 0.7618, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 4.28994083404541, + "learning_rate": 2e-05, + "loss": 0.6162, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.2983256578445435, + "learning_rate": 2e-05, + "loss": 0.1813, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.0821826457977295, + "learning_rate": 2e-05, + "loss": 0.1131, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.355536460876465, + "learning_rate": 2e-05, + "loss": 0.5493, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.08770351856946945, + "learning_rate": 2e-05, + "loss": 0.0152, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.755689024925232, + "learning_rate": 2e-05, + "loss": 0.2324, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.27326440811157227, + "learning_rate": 2e-05, + "loss": 0.3221, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.9416086673736572, + "learning_rate": 2e-05, + "loss": 0.2913, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5324358893436928.0, + "train_loss": 0.2760009717941284, + "train_runtime": 137.4214, + "train_samples_per_second": 2.911, + "train_steps_per_second": 0.728 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5324358893436928.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..964fa249f56fe6d148697088e87e06d9a4c6042a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:53272380bac4d669e1eda5efead2a664e966e4c032a88db582a624b107faa96d +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..24c693fccc026ad11f3af626de4a8f83c745b5a2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:900739f390ed6c00ed98cecfd0741082ec537567df10038a376247067be3afa3 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e87c369d14b573a4f8b03820d66af58f94fceabd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:880ee2be8b15f137e3575e1c23a6fffcefd2797724ad7c011b58674b7f08599b +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..9830ebcfb64b31f817ad36deedd03dda7f55ebc5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3ecafff29c689f87878f30baca71263e8977812f841f21e4cabfbba2649649d9 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..a485b970c0d964004591373df98229206f6ef395 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fa627dc2251a126ca60c518eba68f85a51f88cbdbe20f24b02b7c5162eb13d49 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..12d3561aec15f5631de1622ff9b7c3eae4c45dad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ae678d76844e5aa1b6e60961f0402352778fad2661edf020531ea7cbfeb36979 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..d1943d701dd1703517bb4a9b8b05f4c4771c2a99 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6172bdf1f4d96cb4d1b98fbd58f1968b23503dfd0af420e3aba5d1a23f2ac2a5 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..a19dd79f527736a753ec60dc33ea2253c52960cf --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47283c46b1cc289c7d07e523b017c830d3534b17c8e4fa2c1444beba39aa5d28 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..599702af300999e2445f1bee699d75a0f6e6580e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/14_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.6418522000312805, + "learning_rate": 2e-05, + "loss": 0.449, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.284202814102173, + "learning_rate": 2e-05, + "loss": 0.2883, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.17431582510471344, + "learning_rate": 2e-05, + "loss": 0.0583, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.07775133848190308, + "learning_rate": 2e-05, + "loss": 0.0018, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 6.088167667388916, + "learning_rate": 2e-05, + "loss": 0.6732, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.0427672863006592, + "learning_rate": 2e-05, + "loss": 0.2353, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 3.381765604019165, + "learning_rate": 2e-05, + "loss": 0.3251, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.02370993234217167, + "learning_rate": 2e-05, + "loss": 0.0041, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 3.1150500774383545, + "learning_rate": 2e-05, + "loss": 0.124, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.3724212050437927, + "learning_rate": 2e-05, + "loss": 0.0288, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.1521490514278412, + "learning_rate": 2e-05, + "loss": 0.0248, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.02621975541114807, + "learning_rate": 2e-05, + "loss": 0.0017, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 6.285457611083984, + "learning_rate": 2e-05, + "loss": 0.6485, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 6.891001224517822, + "learning_rate": 2e-05, + "loss": 0.3212, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.296241044998169, + "learning_rate": 2e-05, + "loss": 0.2189, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.8597345948219299, + "learning_rate": 2e-05, + "loss": 0.0739, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.1707541048526764, + "learning_rate": 2e-05, + "loss": 0.1632, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.6294689774513245, + "learning_rate": 2e-05, + "loss": 0.035, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.1435924470424652, + "learning_rate": 2e-05, + "loss": 0.0151, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.5809268951416016, + "learning_rate": 2e-05, + "loss": 0.3765, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.20538181066513062, + "learning_rate": 2e-05, + "loss": 0.1409, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.3682114779949188, + "learning_rate": 2e-05, + "loss": 0.4047, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4620591700077057, + "learning_rate": 2e-05, + "loss": 0.0357, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.1595410108566284, + "learning_rate": 2e-05, + "loss": 0.0663, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.029500562697649002, + "learning_rate": 2e-05, + "loss": 0.172, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.03507590293884277, + "learning_rate": 2e-05, + "loss": 0.0023, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.76004159450531, + "learning_rate": 2e-05, + "loss": 0.0591, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5642727613449097, + "learning_rate": 2e-05, + "loss": 0.0933, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.1975655555725098, + "learning_rate": 2e-05, + "loss": 0.2115, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.013108850456774235, + "learning_rate": 2e-05, + "loss": 0.0033, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.6527280211448669, + "learning_rate": 2e-05, + "loss": 0.2406, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.35223186016082764, + "learning_rate": 2e-05, + "loss": 0.0247, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.676079511642456, + "learning_rate": 2e-05, + "loss": 0.1003, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.5970451831817627, + "learning_rate": 2e-05, + "loss": 0.191, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.4674850106239319, + "learning_rate": 2e-05, + "loss": 0.1579, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.2515055537223816, + "learning_rate": 2e-05, + "loss": 0.0306, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.6629464626312256, + "learning_rate": 2e-05, + "loss": 0.1352, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.2369930744171143, + "learning_rate": 2e-05, + "loss": 0.1498, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.6930259466171265, + "learning_rate": 2e-05, + "loss": 0.105, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.082872152328491, + "learning_rate": 2e-05, + "loss": 0.0706, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.45335134863853455, + "learning_rate": 2e-05, + "loss": 0.0505, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.747994065284729, + "learning_rate": 2e-05, + "loss": 0.0613, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.5007660984992981, + "learning_rate": 2e-05, + "loss": 0.0208, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.2974320352077484, + "learning_rate": 2e-05, + "loss": 0.0221, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.1478660106658936, + "learning_rate": 2e-05, + "loss": 0.0946, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.2641757726669312, + "learning_rate": 2e-05, + "loss": 0.0292, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.10433924198150635, + "learning_rate": 2e-05, + "loss": 0.0156, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.3311959505081177, + "learning_rate": 2e-05, + "loss": 0.0282, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.2857125699520111, + "learning_rate": 2e-05, + "loss": 0.0126, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.004490447696298361, + "learning_rate": 2e-05, + "loss": 0.0318, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5339656207990784.0, + "train_loss": 0.13656323671340942, + "train_runtime": 129.0651, + "train_samples_per_second": 3.099, + "train_steps_per_second": 0.775 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5339656207990784.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..ea41006d7d69f6285650151e9155352288c7516f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d819db2820a5263d5a559c037b85d32c5d01eba53e3b1bfb267ecb172a5566ac +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..114e834bcb1ca85b0d9e0803ab0d2990898bdc5f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8e54ba7f395b8c8bdfe2ea66fa1d9172a1e433a4e299a6ea4d917bf24eaef634 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e434eccb3cd922c213d623ec27f924cf78e442b3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:efc23b36d5f9f37bac728912a9ac690d9ff80b9361a54461b8297638df59becc +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..5fabd69c182427e3933d10cdcc752b118e38bdc6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c6174d1dd9ae3b92994a640afecb186a1faa5d590b7dbda0bdd34f251564cd26 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..04b568bceba98664d04414f7e6bc5689d75aa297 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d15bff61987c52ed0efe5ee65dda0fff39e1ba22604a656663ae2ce338bed998 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..ce477784ebdcbb8ad07eafc4676bee052e6f30db --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1e958cf5a34fa4a50de1855bdbbfb62a25b2eba8bd06671de51d285deb1ce373 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..a7b34f5d737f6b57f6d5a492c97e965eca8dfd96 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f17872b25ad13fbb0d913e108eeaf6434a7a83d4bd141bcbdd8d03b6684adf59 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..b08da4a42992a70022e7bee2da09014d4bcb8017 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3acf43a69e26a8399c9fa756eedb377c37b3f6b1bd0b5637e4982a9600932e24 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..f20e13e3e335982d6f08102d3de4f9c6100d7b00 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/15_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.679689884185791, + "learning_rate": 2e-05, + "loss": 0.1574, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 4.732466697692871, + "learning_rate": 2e-05, + "loss": 0.3729, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 4.097064018249512, + "learning_rate": 2e-05, + "loss": 0.3976, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.246448278427124, + "learning_rate": 2e-05, + "loss": 0.0864, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.7670771479606628, + "learning_rate": 2e-05, + "loss": 0.0771, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.8492998480796814, + "learning_rate": 2e-05, + "loss": 0.1821, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.7263357639312744, + "learning_rate": 2e-05, + "loss": 0.0396, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.5115649104118347, + "learning_rate": 2e-05, + "loss": 0.3179, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.4502179622650146, + "learning_rate": 2e-05, + "loss": 0.0843, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.181198835372925, + "learning_rate": 2e-05, + "loss": 0.1755, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 5.239195823669434, + "learning_rate": 2e-05, + "loss": 0.4381, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.8284292221069336, + "learning_rate": 2e-05, + "loss": 0.2066, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 3.3456647396087646, + "learning_rate": 2e-05, + "loss": 0.3615, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.5233908891677856, + "learning_rate": 2e-05, + "loss": 0.2677, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.5443874597549438, + "learning_rate": 2e-05, + "loss": 0.2126, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.4946682453155518, + "learning_rate": 2e-05, + "loss": 0.0662, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.9547256231307983, + "learning_rate": 2e-05, + "loss": 0.4345, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.14249467849731445, + "learning_rate": 2e-05, + "loss": 0.0444, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.069937229156494, + "learning_rate": 2e-05, + "loss": 0.2662, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.2425663471221924, + "learning_rate": 2e-05, + "loss": 0.1427, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.21746715903282166, + "learning_rate": 2e-05, + "loss": 0.0681, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.425785541534424, + "learning_rate": 2e-05, + "loss": 0.5885, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.5294876098632812, + "learning_rate": 2e-05, + "loss": 0.0556, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 2.124783992767334, + "learning_rate": 2e-05, + "loss": 0.0836, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.257133483886719, + "learning_rate": 2e-05, + "loss": 0.6421, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.9828795194625854, + "learning_rate": 2e-05, + "loss": 0.2599, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.10798734426498413, + "learning_rate": 2e-05, + "loss": 0.0136, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4329739809036255, + "learning_rate": 2e-05, + "loss": 0.0308, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.117771863937378, + "learning_rate": 2e-05, + "loss": 0.2156, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5787556767463684, + "learning_rate": 2e-05, + "loss": 0.0929, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 8.790836334228516, + "learning_rate": 2e-05, + "loss": 0.5723, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.573124408721924, + "learning_rate": 2e-05, + "loss": 0.1222, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.3422446250915527, + "learning_rate": 2e-05, + "loss": 0.1791, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 6.464678764343262, + "learning_rate": 2e-05, + "loss": 0.5707, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.49334731698036194, + "learning_rate": 2e-05, + "loss": 0.3891, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.468083381652832, + "learning_rate": 2e-05, + "loss": 0.2144, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.39885401725769043, + "learning_rate": 2e-05, + "loss": 0.1513, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.8221912384033203, + "learning_rate": 2e-05, + "loss": 0.3558, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.9635553359985352, + "learning_rate": 2e-05, + "loss": 0.0444, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.3474299907684326, + "learning_rate": 2e-05, + "loss": 0.158, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.721296787261963, + "learning_rate": 2e-05, + "loss": 0.1708, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.414232611656189, + "learning_rate": 2e-05, + "loss": 0.167, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.8499486446380615, + "learning_rate": 2e-05, + "loss": 0.5853, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 4.713083267211914, + "learning_rate": 2e-05, + "loss": 0.6089, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.136268138885498, + "learning_rate": 2e-05, + "loss": 0.2493, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.28218191862106323, + "learning_rate": 2e-05, + "loss": 0.1814, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.3688507080078125, + "learning_rate": 2e-05, + "loss": 0.0431, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.38892462849617004, + "learning_rate": 2e-05, + "loss": 0.7692, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 3.604526996612549, + "learning_rate": 2e-05, + "loss": 0.2717, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.6390445232391357, + "learning_rate": 2e-05, + "loss": 0.104, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2214526502043648.0, + "train_loss": 0.2458007049560547, + "train_runtime": 87.3087, + "train_samples_per_second": 4.581, + "train_steps_per_second": 1.145 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2214526502043648.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..d3857c5eafc59f50efa2df28ff09ebeae1cf3f28 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:71f6a7820b89d111811ecda4d1eff2efe5116256076ba2528dd126f010679e43 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..9c4ccf637ad2fce0f96d58df809b412d38cc1c68 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f196670f04ee010a0c38139785209e028ad430e5ce1c6961230aeabdf0deb755 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..53b08952383234d352fcf94790a84bdf6dd7e66a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:02207c03eba12a46a576afa847fe8d5f576e4b6a546bbfa8f4d0c528779df55f +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..80d0bdacfe98885e6b43e6c2a688b3974108dac6 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2a17cbd37e524e66aca044265ff02066f289a88f6502a5ca8fb6ead2a7c9c9c5 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..55dbcebf332a4bb1e97863c072f3814357dda1bd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6fa4b2e1784d7d4999b26a93a1bc839a4c58e87a37eac34a5a4122cc70b669b1 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..e77f2d241412de69edf00583be9a1f3fd27f848f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c186d0f35855ee62490c47917d0a90fd46e9ca8538d89f8d118579ac519062f7 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..2e1b2ffb65562aff0e7c9c312b60c2a2b06b55d4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:11c36d4eab7ac014a98bf95e7682517fcb80b72345be20621ef6fe85dda9ef13 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..71f484d3c25739d950e84eef4d694f62be2631ed --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9b669f6e6dc445dad5909e39d430cba739904d723cc4cc6a25d146e00ebd6826 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..454ec86b1d1a370d8b3fe4131354fc18de74c0dd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/16_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.3187952041625977, + "learning_rate": 2e-05, + "loss": 0.2583, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.24948912858963013, + "learning_rate": 2e-05, + "loss": 0.0301, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 9.049494743347168, + "learning_rate": 2e-05, + "loss": 0.8557, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 4.179375171661377, + "learning_rate": 2e-05, + "loss": 0.2056, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.5302295684814453, + "learning_rate": 2e-05, + "loss": 0.1946, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.0959894061088562, + "learning_rate": 2e-05, + "loss": 0.0171, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.6924583911895752, + "learning_rate": 2e-05, + "loss": 0.1305, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.918614625930786, + "learning_rate": 2e-05, + "loss": 0.3255, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.299540638923645, + "learning_rate": 2e-05, + "loss": 0.0567, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.07549870014190674, + "learning_rate": 2e-05, + "loss": 0.0155, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 4.870273113250732, + "learning_rate": 2e-05, + "loss": 0.6287, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.4728110134601593, + "learning_rate": 2e-05, + "loss": 0.1823, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.5354630947113037, + "learning_rate": 2e-05, + "loss": 0.5062, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.7117526531219482, + "learning_rate": 2e-05, + "loss": 0.1667, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.1585791110992432, + "learning_rate": 2e-05, + "loss": 0.5229, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 3.9950287342071533, + "learning_rate": 2e-05, + "loss": 0.2529, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.28591057658195496, + "learning_rate": 2e-05, + "loss": 0.2372, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.9526095390319824, + "learning_rate": 2e-05, + "loss": 0.2253, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.4037237167358398, + "learning_rate": 2e-05, + "loss": 0.0976, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.828270196914673, + "learning_rate": 2e-05, + "loss": 0.2756, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.0896509513258934, + "learning_rate": 2e-05, + "loss": 0.0068, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.4235999286174774, + "learning_rate": 2e-05, + "loss": 0.1004, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.7319014072418213, + "learning_rate": 2e-05, + "loss": 0.0896, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.7521741390228271, + "learning_rate": 2e-05, + "loss": 0.1007, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.25479304790496826, + "learning_rate": 2e-05, + "loss": 0.0408, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.7026548385620117, + "learning_rate": 2e-05, + "loss": 0.0492, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.3204265832901, + "learning_rate": 2e-05, + "loss": 0.0567, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 5.341701507568359, + "learning_rate": 2e-05, + "loss": 0.3901, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.9164679050445557, + "learning_rate": 2e-05, + "loss": 0.2219, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8323821425437927, + "learning_rate": 2e-05, + "loss": 0.0683, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.42628180980682373, + "learning_rate": 2e-05, + "loss": 0.0195, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.7549359798431396, + "learning_rate": 2e-05, + "loss": 0.4599, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.488187313079834, + "learning_rate": 2e-05, + "loss": 0.0882, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 8.973003387451172, + "learning_rate": 2e-05, + "loss": 1.4797, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.6592448949813843, + "learning_rate": 2e-05, + "loss": 0.1721, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 6.294589996337891, + "learning_rate": 2e-05, + "loss": 0.313, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 5.151583194732666, + "learning_rate": 2e-05, + "loss": 0.731, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.03851696476340294, + "learning_rate": 2e-05, + "loss": 0.0039, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.4949426651000977, + "learning_rate": 2e-05, + "loss": 0.0383, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.31677430868148804, + "learning_rate": 2e-05, + "loss": 0.0374, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.353370428085327, + "learning_rate": 2e-05, + "loss": 0.4504, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.1081212759017944, + "learning_rate": 2e-05, + "loss": 0.2276, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.8206987380981445, + "learning_rate": 2e-05, + "loss": 0.0455, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 3.7010061740875244, + "learning_rate": 2e-05, + "loss": 0.2457, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.9961206912994385, + "learning_rate": 2e-05, + "loss": 0.4304, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.5777392387390137, + "learning_rate": 2e-05, + "loss": 0.1151, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.974031925201416, + "learning_rate": 2e-05, + "loss": 0.7548, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 5.044471740722656, + "learning_rate": 2e-05, + "loss": 0.1583, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.661238193511963, + "learning_rate": 2e-05, + "loss": 0.0914, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 2.1589114665985107, + "learning_rate": 2e-05, + "loss": 0.1359, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2207748649385984.0, + "train_loss": 0.24615520477294922, + "train_runtime": 83.4459, + "train_samples_per_second": 4.794, + "train_steps_per_second": 1.198 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2207748649385984.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..9f56337fa236816dce30d83440b376183d3fd898 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a5bccfdd6561481329c02e5f63bf6f8659b9c4db031c20dd34df146c827e10ef +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..0d37f46c2dd122daeed9d8d37b9b20c3ccb7b25d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0633c2ac165a14945804d4e34851082c2f0fc84a03f16dd266092dd5b3c1f22f +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e10f9ad3194d0a6a381b8ed53d10f2d69835f222 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be698346f0474bcd4a6eaa3ad1830b5dc72b25558a2542662133257922a75832 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..78e99a7ede3c8b141e1279b73482fa719bed43c5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:82afed0b26485687c0cdfec5458470268b88749a7d97cababa32fe164268cfcd +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..1bcc6209c914e4e14c89a8763736fa88fd3c0515 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2b52f019e1643194032bd7da85ee21ee048218c804aa8186129918d320a73d1 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..4a924dd1229189ffd07ccbff43ffe90629fc8f0b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b797c5a07f356021d1c1ade5599ce0899212e3ed2914ff8d7801d64ee250776e +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..8ad8cdfbe5a7efcf43a27149e95fdb669a4dbae3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:74002eed61cdb50902d2e565679d794d1e9ce1c9d105746ac0855de24a57ef0f +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..c4e71d9c488e88d11dbb30dd2fde219b98344248 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:586cd99bf9875e5ef2ff70e63f29b26ad36f2a9f056bfe80599d8df28a7b4dc9 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..bcaf13d4b8c0198c215e8692fc3d10c59bb20e83 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/17_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.9268475770950317, + "learning_rate": 2e-05, + "loss": 0.0616, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.5034847259521484, + "learning_rate": 2e-05, + "loss": 0.4217, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.2765886783599854, + "learning_rate": 2e-05, + "loss": 0.1259, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.7694688439369202, + "learning_rate": 2e-05, + "loss": 0.0518, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 4.505555629730225, + "learning_rate": 2e-05, + "loss": 0.15, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.6297736167907715, + "learning_rate": 2e-05, + "loss": 0.197, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.7030960321426392, + "learning_rate": 2e-05, + "loss": 0.0244, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 5.285088062286377, + "learning_rate": 2e-05, + "loss": 0.2457, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 7.324528694152832, + "learning_rate": 2e-05, + "loss": 0.4305, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 7.067599296569824, + "learning_rate": 2e-05, + "loss": 0.9352, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.351841926574707, + "learning_rate": 2e-05, + "loss": 0.1628, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 6.422041893005371, + "learning_rate": 2e-05, + "loss": 0.7145, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.9146904349327087, + "learning_rate": 2e-05, + "loss": 0.0915, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.15354321897029877, + "learning_rate": 2e-05, + "loss": 0.0095, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.5202783346176147, + "learning_rate": 2e-05, + "loss": 0.129, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.3365912437438965, + "learning_rate": 2e-05, + "loss": 0.1203, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8790549635887146, + "learning_rate": 2e-05, + "loss": 0.0656, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.3751930296421051, + "learning_rate": 2e-05, + "loss": 0.0179, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 2.5937938690185547, + "learning_rate": 2e-05, + "loss": 0.2616, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.4512746334075928, + "learning_rate": 2e-05, + "loss": 0.2855, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.3875581622123718, + "learning_rate": 2e-05, + "loss": 0.0224, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.9946991205215454, + "learning_rate": 2e-05, + "loss": 0.2919, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.271794319152832, + "learning_rate": 2e-05, + "loss": 0.075, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 4.7678327560424805, + "learning_rate": 2e-05, + "loss": 0.2516, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.7908960580825806, + "learning_rate": 2e-05, + "loss": 0.371, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.192012071609497, + "learning_rate": 2e-05, + "loss": 0.1888, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 2.0684874057769775, + "learning_rate": 2e-05, + "loss": 0.1676, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 4.743781566619873, + "learning_rate": 2e-05, + "loss": 0.2756, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 5.399972438812256, + "learning_rate": 2e-05, + "loss": 0.3251, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 6.208322525024414, + "learning_rate": 2e-05, + "loss": 0.4799, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 7.416908264160156, + "learning_rate": 2e-05, + "loss": 1.0367, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 4.551591396331787, + "learning_rate": 2e-05, + "loss": 0.8719, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.263587087392807, + "learning_rate": 2e-05, + "loss": 0.0621, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 5.655372619628906, + "learning_rate": 2e-05, + "loss": 0.2745, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.014061689376831, + "learning_rate": 2e-05, + "loss": 0.049, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.2437562942504883, + "learning_rate": 2e-05, + "loss": 0.4927, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.776449203491211, + "learning_rate": 2e-05, + "loss": 0.3078, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 3.584463357925415, + "learning_rate": 2e-05, + "loss": 0.6783, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 2.419666290283203, + "learning_rate": 2e-05, + "loss": 0.1796, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.34660226106643677, + "learning_rate": 2e-05, + "loss": 0.1267, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 2.910494089126587, + "learning_rate": 2e-05, + "loss": 0.2945, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.8461408615112305, + "learning_rate": 2e-05, + "loss": 0.4319, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.7095017433166504, + "learning_rate": 2e-05, + "loss": 0.3939, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.085622787475586, + "learning_rate": 2e-05, + "loss": 0.0401, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 3.8005211353302, + "learning_rate": 2e-05, + "loss": 0.5405, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.445932626724243, + "learning_rate": 2e-05, + "loss": 0.1669, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.0819096565246582, + "learning_rate": 2e-05, + "loss": 0.0493, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.9659722447395325, + "learning_rate": 2e-05, + "loss": 0.0515, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.038895342499017715, + "learning_rate": 2e-05, + "loss": 0.0884, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.154855728149414, + "learning_rate": 2e-05, + "loss": 0.0617, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2212818824724480.0, + "train_loss": 0.26297718048095703, + "train_runtime": 81.406, + "train_samples_per_second": 4.914, + "train_steps_per_second": 1.228 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2212818824724480.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..8d8e70c0a2c9413e8ca60323ec6ff62da980a203 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2d9f33c10f40f6d892e0ad3bde71b8c631dd28121b69115de08bfa87cf432399 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..85880e90606f6dc6c9211f89521e3acb52944df5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cff5fe10422983275c67a87a972a1dcaa607f807d3d5156d515d25a6e3408f0c +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..a9c726be0c5cbc02008a904ecf36e32d447a5ab1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:df8b43f2837cbe071bda1a67d7d7c360c62e5cf399e78b51e4f31ebe9e0a3329 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..2fdb3e343e77fffd659748d12a95aca2f738286d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b3a953b15b35b514f54d7143d0a93bcd416e539656309a670cd218e7b05f313 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..2925ce85b79f99f2f3a22baaa9285199ce27fb90 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:67b06fb225216f4a65567b1482ed7164427b69b185807b619302243bb7221ca9 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..3e788ae3bcfb7b10562c776777727ad60cda9697 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6a5693265164ef2fd9a965ff32fbc3d987b3de7c63df1ddf79b4ee1e1755b9c7 +size 369839594 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..0c3432d174d13c1188c39ca92fb8c5cdca086684 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72d3df3ee9e8248485d1bad878e36697058abe55cf4dae3002c53b17a0e3fe47 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..cb24d2bcd3900221cc2e472cba01e49819653ee4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6ab10a9358ec7701f99d2933e946404e8e5390aca38bb858e25dacfd4e2b71f9 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..7f8e8ada5363eef10dddbf0b1fbfaa1b3b46f346 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/18_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.07982902228832245, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.2455527782440186, + "learning_rate": 2e-05, + "loss": 0.0695, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.4779789447784424, + "learning_rate": 2e-05, + "loss": 0.0154, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.043753936886787415, + "learning_rate": 2e-05, + "loss": 0.002, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.06418295949697495, + "learning_rate": 2e-05, + "loss": 0.45, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 5.823064804077148, + "learning_rate": 2e-05, + "loss": 0.1803, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.6937955021858215, + "learning_rate": 2e-05, + "loss": 0.138, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.23485641181468964, + "learning_rate": 2e-05, + "loss": 0.0424, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.0868310928344727, + "learning_rate": 2e-05, + "loss": 0.0461, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.597323179244995, + "learning_rate": 2e-05, + "loss": 0.1981, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.07463926076889038, + "learning_rate": 2e-05, + "loss": 0.1084, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.3387567400932312, + "learning_rate": 2e-05, + "loss": 0.1006, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.408559650182724, + "learning_rate": 2e-05, + "loss": 0.0774, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 4.313039779663086, + "learning_rate": 2e-05, + "loss": 0.4198, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.04163407161831856, + "learning_rate": 2e-05, + "loss": 0.0424, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 5.253049850463867, + "learning_rate": 2e-05, + "loss": 0.2253, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5455938577651978, + "learning_rate": 2e-05, + "loss": 0.0348, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.201729878783226, + "learning_rate": 2e-05, + "loss": 0.1964, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.3942170143127441, + "learning_rate": 2e-05, + "loss": 0.0649, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.2602633535861969, + "learning_rate": 2e-05, + "loss": 0.028, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.7189749479293823, + "learning_rate": 2e-05, + "loss": 0.3457, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.14480771124362946, + "learning_rate": 2e-05, + "loss": 0.0071, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.7231111526489258, + "learning_rate": 2e-05, + "loss": 0.3886, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.513663411140442, + "learning_rate": 2e-05, + "loss": 0.0869, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.1579140424728394, + "learning_rate": 2e-05, + "loss": 0.1049, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.8974907398223877, + "learning_rate": 2e-05, + "loss": 0.1034, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 3.530357837677002, + "learning_rate": 2e-05, + "loss": 0.1745, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.28248777985572815, + "learning_rate": 2e-05, + "loss": 0.3563, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.1568111777305603, + "learning_rate": 2e-05, + "loss": 0.0542, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.801914632320404, + "learning_rate": 2e-05, + "loss": 0.0718, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.006407019216567278, + "learning_rate": 2e-05, + "loss": 0.0228, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.0029070377349854, + "learning_rate": 2e-05, + "loss": 0.1002, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.04014541208744049, + "learning_rate": 2e-05, + "loss": 0.1406, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.0483016967773438, + "learning_rate": 2e-05, + "loss": 0.0721, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 5.478046417236328, + "learning_rate": 2e-05, + "loss": 0.4529, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.08521467447280884, + "learning_rate": 2e-05, + "loss": 0.0143, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.9899935722351074, + "learning_rate": 2e-05, + "loss": 0.0981, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.932433128356934, + "learning_rate": 2e-05, + "loss": 0.1511, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 4.54806661605835, + "learning_rate": 2e-05, + "loss": 0.2866, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.27968883514404297, + "learning_rate": 2e-05, + "loss": 0.0095, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.8157901763916016, + "learning_rate": 2e-05, + "loss": 0.2934, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.09253838658332825, + "learning_rate": 2e-05, + "loss": 0.1175, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.6851876378059387, + "learning_rate": 2e-05, + "loss": 0.0177, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.10146772861480713, + "learning_rate": 2e-05, + "loss": 0.014, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 2.075961112976074, + "learning_rate": 2e-05, + "loss": 0.1398, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.4204301834106445, + "learning_rate": 2e-05, + "loss": 0.1513, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 5.38856315612793, + "learning_rate": 2e-05, + "loss": 0.3948, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 5.210330009460449, + "learning_rate": 2e-05, + "loss": 0.1038, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.02752024680376053, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.9672114849090576, + "learning_rate": 2e-05, + "loss": 0.2289, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2207554989981696.0, + "train_loss": 0.1389507007598877, + "train_runtime": 80.5198, + "train_samples_per_second": 4.968, + "train_steps_per_second": 1.242 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2207554989981696.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..233d634734322f63b2b1172b515877e69e906dda --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dd651629ac4b48404f1761ccec753c09fdd8883435f2b9c11f63385cef243552 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..b9a974a171835794cfe9347a14352539292c49d7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7cb9804477fd50bcef834bd1311e60d34734ed2abf8a1864e2e921caacd5283d +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..069458914ef4b90abb6b1fd9e2cf479f5757e5fd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:42e3695be0bbf6d9cb494a7e672f38fadab7147df09e7576211630843e607cb1 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..5d81c144503f7a7bae5bcc900593f4200aa74c66 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63e3995d1bfe84f1656754bafe6e25d22d8bd643ac0eb6ce58221f0850db3abc +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..072345a8f850169c433aa6fd99b60b2c0bb92cdb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:46dfb9d4e7928a9e8a9b00b359c17d2e2f06fc615545d8f3fbe707a3d81f3fca +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..6323d807e296c7ddd6f3f4d7efd3f3e947595071 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:105d4fff914fb529fb39429d8804d16d16ff7ce12d22451b3c477ab6e84974db +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..03c66d2f51915b055943d921df55e43b40962a83 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92ab7f5f5ebe3d322f0c3238f805aa012799c894c7244b9314e1489a57a2988a +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..7548f3119ff2321330de12be14b7159abcbaad48 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8b094064408c32f3cfd4a4dc31b83c3f734d6af98623509e253d5ec94c05a643 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..229cf6f4fba324c18833b0ea5f3b16c6948555e1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/19_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.22729668021202087, + "learning_rate": 2e-05, + "loss": 0.0867, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.31728029251098633, + "learning_rate": 2e-05, + "loss": 0.0718, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.45670491456985474, + "learning_rate": 2e-05, + "loss": 0.0397, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.6296160221099854, + "learning_rate": 2e-05, + "loss": 0.1279, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.585667610168457, + "learning_rate": 2e-05, + "loss": 0.3575, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.23326006531715393, + "learning_rate": 2e-05, + "loss": 0.047, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.7845252752304077, + "learning_rate": 2e-05, + "loss": 0.0608, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.4553896188735962, + "learning_rate": 2e-05, + "loss": 0.1141, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.4276368319988251, + "learning_rate": 2e-05, + "loss": 0.2497, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.10590571910142899, + "learning_rate": 2e-05, + "loss": 0.0104, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.9433478713035583, + "learning_rate": 2e-05, + "loss": 0.1149, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.9218755960464478, + "learning_rate": 2e-05, + "loss": 0.0633, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.5408377647399902, + "learning_rate": 2e-05, + "loss": 0.135, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.2340569496154785, + "learning_rate": 2e-05, + "loss": 0.1252, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 2.2514712810516357, + "learning_rate": 2e-05, + "loss": 0.2737, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.39434510469436646, + "learning_rate": 2e-05, + "loss": 0.4539, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.8574960231781006, + "learning_rate": 2e-05, + "loss": 0.2067, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.42364007234573364, + "learning_rate": 2e-05, + "loss": 0.1053, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.18779227137565613, + "learning_rate": 2e-05, + "loss": 0.0359, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.07692208141088486, + "learning_rate": 2e-05, + "loss": 0.065, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.7394633889198303, + "learning_rate": 2e-05, + "loss": 0.0588, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.5901036262512207, + "learning_rate": 2e-05, + "loss": 0.5377, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4624583423137665, + "learning_rate": 2e-05, + "loss": 0.1289, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.25090086460113525, + "learning_rate": 2e-05, + "loss": 0.2409, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.21005111932754517, + "learning_rate": 2e-05, + "loss": 0.2444, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.6578717231750488, + "learning_rate": 2e-05, + "loss": 0.2062, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.2951827943325043, + "learning_rate": 2e-05, + "loss": 0.1687, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.11019435524940491, + "learning_rate": 2e-05, + "loss": 0.0121, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.788218379020691, + "learning_rate": 2e-05, + "loss": 0.1742, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.8302671909332275, + "learning_rate": 2e-05, + "loss": 0.0626, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.07048580795526505, + "learning_rate": 2e-05, + "loss": 0.108, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 2.9046387672424316, + "learning_rate": 2e-05, + "loss": 0.2656, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.1957688331604004, + "learning_rate": 2e-05, + "loss": 0.3242, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.06359727680683136, + "learning_rate": 2e-05, + "loss": 0.0788, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.589907646179199, + "learning_rate": 2e-05, + "loss": 0.3987, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.8452297449111938, + "learning_rate": 2e-05, + "loss": 0.1224, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.24079157412052155, + "learning_rate": 2e-05, + "loss": 0.0306, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 4.565906047821045, + "learning_rate": 2e-05, + "loss": 0.4745, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.9080637693405151, + "learning_rate": 2e-05, + "loss": 0.1578, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.1943283081054688, + "learning_rate": 2e-05, + "loss": 0.1306, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.4385234117507935, + "learning_rate": 2e-05, + "loss": 0.3276, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.10782930254936218, + "learning_rate": 2e-05, + "loss": 0.0254, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.001715279184281826, + "learning_rate": 2e-05, + "loss": 0.0319, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.16809244453907013, + "learning_rate": 2e-05, + "loss": 0.2837, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5914433598518372, + "learning_rate": 2e-05, + "loss": 0.0303, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.7493177056312561, + "learning_rate": 2e-05, + "loss": 0.0987, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.40446290373802185, + "learning_rate": 2e-05, + "loss": 0.1004, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 1.3253214359283447, + "learning_rate": 2e-05, + "loss": 0.172, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.13731542229652405, + "learning_rate": 2e-05, + "loss": 0.0345, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.06778030842542648, + "learning_rate": 2e-05, + "loss": 0.2492, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5295755845697536.0, + "train_loss": 0.1604784035682678, + "train_runtime": 125.6432, + "train_samples_per_second": 3.184, + "train_steps_per_second": 0.796 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5295755845697536.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..dca05f8fb10abb9b1b88730d926ee71f7d9c3dcb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2ca597905f53a3fa129b2a01b0fa5b9e79cf05ecb54e14e15d05bad174ae3af8 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..5fbbd53ffec9ef582654a2a8be3a695993d1b775 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c254b0c5f7cac63c51dc50f90bb6706bcb597ac405360a6d33d7109d1cdae7c5 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..131bd4edf1a8359eab5808be767765a6351a4348 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5f16a0ade9bb5185cee2729699c77bb414fd64d2fa118319a2da1cca39045ed4 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..d6b3cd2314ddd7d4db8302956836827d44a5a5df --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a7c347a97cca6d3652cd458852c15985daa8a070df17487148fa7caf06650f2b +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..dc6f122d840986b6b756573f005a2393c0f8fb60 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:563839e68a29c52969d48ae41400f15cd1199319e6b50191ff4588b22baad3df +size 369837282 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..a900511834b45e2a3cd8ab26ff872898a4cdec1d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:38b5917223ee56177e3fbc04fac58beb663aa3b05dccb8b875f2c70a000109a8 +size 369838470 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..df07a350c35b8f8b63b683a2572b620089d73848 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7b3da1ea79ad4e810ec69240b5607a77572d9f56697f1629fce2becd993ec8a7 +size 369837282 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..f740d88ae4e8bff1eda1e662d8cd58ef3dbed3df --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:74c113756e6c3017720ecc4c21bb8f71cc27506fc229d7ea5f27e9c7e00c4cff +size 369837282 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..8d7cf24e4ddafc909613764c4ef030d53a609f89 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/1_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.028881167992949486, + "learning_rate": 2e-05, + "loss": 0.017, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.04178139567375183, + "learning_rate": 2e-05, + "loss": 0.0329, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.1752421110868454, + "learning_rate": 2e-05, + "loss": 0.0043, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.1626032292842865, + "learning_rate": 2e-05, + "loss": 0.008, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.070888452231884, + "learning_rate": 2e-05, + "loss": 0.0014, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.4707482159137726, + "learning_rate": 2e-05, + "loss": 0.0321, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.008385924622416496, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.08869297802448273, + "learning_rate": 2e-05, + "loss": 0.0567, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.06399267911911011, + "learning_rate": 2e-05, + "loss": 0.072, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.004949385765939951, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.0222014170140028, + "learning_rate": 2e-05, + "loss": 0.0173, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.11237145960330963, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.3386573791503906, + "learning_rate": 2e-05, + "loss": 0.0879, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.022797059267759323, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.07221390306949615, + "learning_rate": 2e-05, + "loss": 0.0027, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.1886264681816101, + "learning_rate": 2e-05, + "loss": 0.0525, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.7551073431968689, + "learning_rate": 2e-05, + "loss": 0.0161, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.001585116027854383, + "learning_rate": 2e-05, + "loss": 0.0001, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.05067340284585953, + "learning_rate": 2e-05, + "loss": 0.0012, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.8383181095123291, + "learning_rate": 2e-05, + "loss": 0.0173, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.21666410565376282, + "learning_rate": 2e-05, + "loss": 0.0036, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.30005145072937, + "learning_rate": 2e-05, + "loss": 0.1682, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.0034498891327530146, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.027388641610741615, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.0072450763545930386, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.016917383298277855, + "learning_rate": 2e-05, + "loss": 0.0326, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.5184288024902344, + "learning_rate": 2e-05, + "loss": 0.0402, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.02179768681526184, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.005018856376409531, + "learning_rate": 2e-05, + "loss": 0.0739, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.008803858421742916, + "learning_rate": 2e-05, + "loss": 0.0103, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.09704452753067017, + "learning_rate": 2e-05, + "loss": 0.0015, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 0.013792658224701881, + "learning_rate": 2e-05, + "loss": 0.0088, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.03694701939821243, + "learning_rate": 2e-05, + "loss": 0.0007, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.15428417921066284, + "learning_rate": 2e-05, + "loss": 0.0035, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.0039619700983166695, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.6323522329330444, + "learning_rate": 2e-05, + "loss": 0.0104, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.6308977007865906, + "learning_rate": 2e-05, + "loss": 0.0128, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.009072153829038143, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.273922324180603, + "learning_rate": 2e-05, + "loss": 0.0257, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.004881344269961119, + "learning_rate": 2e-05, + "loss": 0.0011, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.006768243853002787, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.0007242965511977673, + "learning_rate": 2e-05, + "loss": 0.0002, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.02291828952729702, + "learning_rate": 2e-05, + "loss": 0.0885, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.005617976654320955, + "learning_rate": 2e-05, + "loss": 0.0008, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.24594473838806152, + "learning_rate": 2e-05, + "loss": 0.0053, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.985758364200592, + "learning_rate": 2e-05, + "loss": 0.0067, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.009313897229731083, + "learning_rate": 2e-05, + "loss": 0.0006, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.2606780529022217, + "learning_rate": 2e-05, + "loss": 0.0047, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.01463312841951847, + "learning_rate": 2e-05, + "loss": 0.0004, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.015957485884428024, + "learning_rate": 2e-05, + "loss": 0.0003, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 2217994633609216.0, + "train_loss": 0.018576545119285585, + "train_runtime": 82.0446, + "train_samples_per_second": 4.875, + "train_steps_per_second": 1.219 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 2217994633609216.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..eb808b3b06807bb6d6754f0272b3a09345667a68 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:57ba46a2e3733cbb5f944b8f31ca97a5ff293a2e4162d932c8e68226ec253965 +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..1b959a9e20102cbefc36b1bb48d35ca3cc0405ab --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c4bef9f01496efd67a226641e4a6e320d52efe1dc3b3d6bb6663084a0a4d839e +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..2b6238cfd94e14949003c0371fc73bb7ddfbb9c2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:181813fd894efb03e673f143942bd9c413da62e6e48953401ec0797d62fe0b8f +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..75ca861f877703d296ec4cc02620e6ff98751c7b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:88e4b8152dc04e907c6985883e8b3a2e70a40e6e9976f0c0ecb03583da93143c +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..8376666407fadd0b6df3b479d0d4ef9125f51e0a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2e32e751a57e9a6063bf1a374efd66c42b2c91841acd24488d91db75a3ea0467 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..99abe0b94d5db0d1db43a01324b7c6ba55b50c28 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b20da0a0597f6e71d679087b41cc2619e14b6f6745221e0fedd1c1b6b88ec19f +size 794710050 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..3d52fb5c6eb0554311f6b43df2da6738da797015 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4fc182df4576bcaaaabf4e157997a17d6fe4250d61070f522a92a2e8dff75b38 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..eda94fa65b66b29ed636b335909fa1dcf6419f92 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b77b36b57fdea86369b252978ce3edda64931b831c78a44bebf22dde112058da +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..ab57f402fd921834614f19eabed825093795ea9d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/20_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.08531734347343445, + "learning_rate": 2e-05, + "loss": 0.0323, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.0926377773284912, + "learning_rate": 2e-05, + "loss": 0.0607, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 2.3126068115234375, + "learning_rate": 2e-05, + "loss": 0.1127, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.6126227974891663, + "learning_rate": 2e-05, + "loss": 0.0291, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.6330201625823975, + "learning_rate": 2e-05, + "loss": 0.4448, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.10275842249393463, + "learning_rate": 2e-05, + "loss": 0.0077, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.29542553424835205, + "learning_rate": 2e-05, + "loss": 0.0417, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.0568504333496094, + "learning_rate": 2e-05, + "loss": 0.0852, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.1004462242126465, + "learning_rate": 2e-05, + "loss": 0.2101, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.14005912840366364, + "learning_rate": 2e-05, + "loss": 0.1686, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.3769957721233368, + "learning_rate": 2e-05, + "loss": 0.0183, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.17943313717842102, + "learning_rate": 2e-05, + "loss": 0.0165, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.234056830406189, + "learning_rate": 2e-05, + "loss": 0.3191, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.1878470778465271, + "learning_rate": 2e-05, + "loss": 0.0159, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.710300087928772, + "learning_rate": 2e-05, + "loss": 0.0259, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.477067470550537, + "learning_rate": 2e-05, + "loss": 0.2618, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.09527702629566193, + "learning_rate": 2e-05, + "loss": 0.0428, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.35215455293655396, + "learning_rate": 2e-05, + "loss": 0.0371, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.23481200635433197, + "learning_rate": 2e-05, + "loss": 0.0701, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 0.5495511889457703, + "learning_rate": 2e-05, + "loss": 0.0724, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.270667165517807, + "learning_rate": 2e-05, + "loss": 0.1419, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.8299638628959656, + "learning_rate": 2e-05, + "loss": 0.0701, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.6716135740280151, + "learning_rate": 2e-05, + "loss": 0.0998, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.557615339756012, + "learning_rate": 2e-05, + "loss": 0.0395, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.1025652885437012, + "learning_rate": 2e-05, + "loss": 0.1521, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.3748072385787964, + "learning_rate": 2e-05, + "loss": 0.0171, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.3879741132259369, + "learning_rate": 2e-05, + "loss": 0.1449, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.05431313440203667, + "learning_rate": 2e-05, + "loss": 0.0323, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.014691580086946487, + "learning_rate": 2e-05, + "loss": 0.1797, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.11475786566734314, + "learning_rate": 2e-05, + "loss": 0.0529, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.01338791847229, + "learning_rate": 2e-05, + "loss": 0.2838, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.827707052230835, + "learning_rate": 2e-05, + "loss": 0.5262, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 5.0494890213012695, + "learning_rate": 2e-05, + "loss": 0.7482, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 3.8526158332824707, + "learning_rate": 2e-05, + "loss": 0.1763, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.6057355403900146, + "learning_rate": 2e-05, + "loss": 0.0568, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.04486304149031639, + "learning_rate": 2e-05, + "loss": 0.0261, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.14988061785697937, + "learning_rate": 2e-05, + "loss": 0.0149, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 5.294212818145752, + "learning_rate": 2e-05, + "loss": 0.5794, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.32698217034339905, + "learning_rate": 2e-05, + "loss": 0.0369, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.013690175488591194, + "learning_rate": 2e-05, + "loss": 0.024, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.09842357784509659, + "learning_rate": 2e-05, + "loss": 0.0672, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.9954636693000793, + "learning_rate": 2e-05, + "loss": 0.1284, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.2176171988248825, + "learning_rate": 2e-05, + "loss": 0.022, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.7552189826965332, + "learning_rate": 2e-05, + "loss": 0.1561, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5929461717605591, + "learning_rate": 2e-05, + "loss": 0.119, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.7603516578674316, + "learning_rate": 2e-05, + "loss": 0.3243, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.6126036047935486, + "learning_rate": 2e-05, + "loss": 0.0401, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.06367256492376328, + "learning_rate": 2e-05, + "loss": 0.0038, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.2796757221221924, + "learning_rate": 2e-05, + "loss": 0.1605, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.07488951832056046, + "learning_rate": 2e-05, + "loss": 0.0069, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5296221983866880.0, + "train_loss": 0.13007930040359497, + "train_runtime": 128.4461, + "train_samples_per_second": 3.114, + "train_steps_per_second": 0.779 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5296221983866880.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..66cb3c827aca139295d577c991b1185ffb1e6263 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f0aeeddf0f0efeaf185d4122529f951dcc1cea785fc515f70db39b2644ed2e2f +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..47a1260be6065aa889bc0c7364a26fda0cad4c74 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8f4c1d7d21e55f812c450f30e2d1b49748aa4ec81d1657a0516523b3a0001b57 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..3cab82a7297acc11c357de073b4e57672cbd829a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:826defa44daa2df97f26dca4c97e76d517872ad79db7f0a7dc35950d82d61529 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..3e10ef644164355c26bdf62fed6d7042fc8088d3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:77010cbdd7412ce934859435130d259450fb84e064be64601e188b33942f7f7e +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..175a4d4d4f6736f7f02bad777ab28872a908d834 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f4d32096cb081ff146b1e31993bbeb56afc6945692a436ea5c3b725acfd81600 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..2910f5dee111bdf2abd9700b1feab5a7d5570e86 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:89db23ef1490ef7a9f476a61a40d5e5507b4b8f2942d14a684223ced3f336a00 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..7740bbcc1a507a79289bcd39b3fb9861b6d44647 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5e35f70ed740743f4a285438fa76e33402591c8f410526e10a8172a42e2e9895 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..e85d26095c7f78ecdc6b8c427afe83e7e436d6b7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b18add1069f9a530f03be0ff4e0dae8e8ada6d70b1b27ccf960c1598940c41c2 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c1776fdbd2d84a848c7a6e6bf35e85a4c2f7f2b1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/2_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.0735808610916138, + "learning_rate": 2e-05, + "loss": 0.1248, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.6208010911941528, + "learning_rate": 2e-05, + "loss": 0.6128, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.2576985359191895, + "learning_rate": 2e-05, + "loss": 0.4083, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.5944550037384033, + "learning_rate": 2e-05, + "loss": 0.3073, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.8347348570823669, + "learning_rate": 2e-05, + "loss": 0.131, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.8249328136444092, + "learning_rate": 2e-05, + "loss": 0.5294, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.22106881439685822, + "learning_rate": 2e-05, + "loss": 0.1678, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.9695217609405518, + "learning_rate": 2e-05, + "loss": 0.2604, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.284844160079956, + "learning_rate": 2e-05, + "loss": 0.2389, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.7450605630874634, + "learning_rate": 2e-05, + "loss": 0.4883, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.2713537812232971, + "learning_rate": 2e-05, + "loss": 0.0648, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.5445595979690552, + "learning_rate": 2e-05, + "loss": 0.2578, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.104795217514038, + "learning_rate": 2e-05, + "loss": 0.5154, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.2727891206741333, + "learning_rate": 2e-05, + "loss": 0.1944, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.2581905424594879, + "learning_rate": 2e-05, + "loss": 0.076, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.1513352394104004, + "learning_rate": 2e-05, + "loss": 0.1873, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.4446638822555542, + "learning_rate": 2e-05, + "loss": 0.3772, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.5228360891342163, + "learning_rate": 2e-05, + "loss": 0.0554, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.2017120122909546, + "learning_rate": 2e-05, + "loss": 0.4803, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.7885148525238037, + "learning_rate": 2e-05, + "loss": 0.4377, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.8897362947463989, + "learning_rate": 2e-05, + "loss": 0.3262, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.3961503207683563, + "learning_rate": 2e-05, + "loss": 0.0732, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.422475814819336, + "learning_rate": 2e-05, + "loss": 0.2614, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.31549346446990967, + "learning_rate": 2e-05, + "loss": 0.1776, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.9473968148231506, + "learning_rate": 2e-05, + "loss": 0.1174, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.5539112091064453, + "learning_rate": 2e-05, + "loss": 0.2704, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.4659912586212158, + "learning_rate": 2e-05, + "loss": 0.2783, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.1413074731826782, + "learning_rate": 2e-05, + "loss": 0.1987, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.9005150198936462, + "learning_rate": 2e-05, + "loss": 0.0673, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.08340349048376083, + "learning_rate": 2e-05, + "loss": 0.1469, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.9565677046775818, + "learning_rate": 2e-05, + "loss": 0.3511, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.3453073501586914, + "learning_rate": 2e-05, + "loss": 0.1838, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.8434649705886841, + "learning_rate": 2e-05, + "loss": 0.1042, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.858950614929199, + "learning_rate": 2e-05, + "loss": 0.2599, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.4917947053909302, + "learning_rate": 2e-05, + "loss": 0.4104, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.6736851334571838, + "learning_rate": 2e-05, + "loss": 0.4978, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.7169210314750671, + "learning_rate": 2e-05, + "loss": 0.0603, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.0006201267242432, + "learning_rate": 2e-05, + "loss": 0.3159, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.8807411789894104, + "learning_rate": 2e-05, + "loss": 0.0827, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.514021396636963, + "learning_rate": 2e-05, + "loss": 0.3767, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.355543375015259, + "learning_rate": 2e-05, + "loss": 1.3726, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.16011284291744232, + "learning_rate": 2e-05, + "loss": 0.0144, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.4076273441314697, + "learning_rate": 2e-05, + "loss": 0.3122, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.4399515390396118, + "learning_rate": 2e-05, + "loss": 0.3164, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 4.456518650054932, + "learning_rate": 2e-05, + "loss": 1.0562, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.226089745759964, + "learning_rate": 2e-05, + "loss": 0.1134, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.032999277114868, + "learning_rate": 2e-05, + "loss": 0.2872, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.3489951193332672, + "learning_rate": 2e-05, + "loss": 0.1211, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.50156307220459, + "learning_rate": 2e-05, + "loss": 0.5985, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.36324140429496765, + "learning_rate": 2e-05, + "loss": 0.0649, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5223464331902976.0, + "train_loss": 0.29464771270751955, + "train_runtime": 127.3644, + "train_samples_per_second": 3.141, + "train_steps_per_second": 0.785 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5223464331902976.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..4581bced1bbeaa8cedd7480ae3ec1320025834d1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:286de596da6de876a8c66a19fa490f54b335c0390bb30659ce202838494a5142 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..43710205356b52e475119506fe0cd39d543abf78 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9282d1dfb9ab5267cafc53c44bfea5884188c36d537fc80029bbd093276477be +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..8fc145a59a316fe7d745693fec5e94141ef97b0d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5438fec1df05bf55fb31cef81cb1cd742107e294af49b9b1aaac454da3b2ace6 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..03052168e2e55d090b5ecde2856cb397aafc879a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7c41b807b65b9d658db504f53bf5aed2ec74512218e89e12d14e31eca8a0fcdd +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..816fc560d9c55f0d221a2b94034545c3b053cced --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e4d6eb9d084b4b4eb0aac853e060bd71aa59765b4bdf158ebc927d7fcc547bcf +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..9a4b74f8a8d5e0512db8c455b817242aa5bd06e5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:53434dbf2977ca9e44cf8b7157ce624eef86b2223e416211e84dd42e5ff45a45 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..632bc45793e5e719d4aed80440f213a175dea144 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a0df7cc67637aadaf6f8a037610011539cf8b270eec4faca52454d74fce4896b +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..2917c5899dfc1f468b86fb18a6f621e0fc9281b5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1590d3cba834eb69e70f49276a752762079d6aa49be4218fa8638071ef125688 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..df4f6a52f88b741c8cf19a9f3a2d5d6fc8fae187 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/3_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 2.5537219047546387, + "learning_rate": 2e-05, + "loss": 0.7915, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.6754408478736877, + "learning_rate": 2e-05, + "loss": 0.3635, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.2286841869354248, + "learning_rate": 2e-05, + "loss": 0.2868, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.4913978576660156, + "learning_rate": 2e-05, + "loss": 0.6548, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 2.601992130279541, + "learning_rate": 2e-05, + "loss": 0.4771, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 2.763522148132324, + "learning_rate": 2e-05, + "loss": 0.7651, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.952918529510498, + "learning_rate": 2e-05, + "loss": 0.729, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.8379621505737305, + "learning_rate": 2e-05, + "loss": 0.5399, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.2027956247329712, + "learning_rate": 2e-05, + "loss": 0.1628, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 0.9438942074775696, + "learning_rate": 2e-05, + "loss": 0.3941, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.6609513759613037, + "learning_rate": 2e-05, + "loss": 0.2372, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 1.4963562488555908, + "learning_rate": 2e-05, + "loss": 0.481, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.3700706958770752, + "learning_rate": 2e-05, + "loss": 0.2188, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.3240225315093994, + "learning_rate": 2e-05, + "loss": 0.3092, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.3214831352233887, + "learning_rate": 2e-05, + "loss": 0.399, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.9053499698638916, + "learning_rate": 2e-05, + "loss": 0.4093, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.6040465831756592, + "learning_rate": 2e-05, + "loss": 0.6965, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.932678699493408, + "learning_rate": 2e-05, + "loss": 0.625, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.2711498737335205, + "learning_rate": 2e-05, + "loss": 0.1973, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.5860810279846191, + "learning_rate": 2e-05, + "loss": 0.4711, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.7776730060577393, + "learning_rate": 2e-05, + "loss": 0.5957, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.3857572078704834, + "learning_rate": 2e-05, + "loss": 0.6201, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.0778429508209229, + "learning_rate": 2e-05, + "loss": 0.2301, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.567577064037323, + "learning_rate": 2e-05, + "loss": 0.0618, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.6822116374969482, + "learning_rate": 2e-05, + "loss": 0.5273, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 3.362592935562134, + "learning_rate": 2e-05, + "loss": 0.7424, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.08990079909563065, + "learning_rate": 2e-05, + "loss": 0.1993, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.1235628128051758, + "learning_rate": 2e-05, + "loss": 0.2549, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.7444652915000916, + "learning_rate": 2e-05, + "loss": 0.2908, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.9353086948394775, + "learning_rate": 2e-05, + "loss": 0.5026, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.116203784942627, + "learning_rate": 2e-05, + "loss": 0.4117, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.6504265069961548, + "learning_rate": 2e-05, + "loss": 0.2925, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.8762640953063965, + "learning_rate": 2e-05, + "loss": 0.5199, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 4.473509311676025, + "learning_rate": 2e-05, + "loss": 0.8141, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.5204918384552, + "learning_rate": 2e-05, + "loss": 0.4309, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.3686333894729614, + "learning_rate": 2e-05, + "loss": 0.6792, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 3.546315908432007, + "learning_rate": 2e-05, + "loss": 0.671, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 1.8682063817977905, + "learning_rate": 2e-05, + "loss": 0.2316, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.6339802742004395, + "learning_rate": 2e-05, + "loss": 0.333, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.4751263856887817, + "learning_rate": 2e-05, + "loss": 0.1728, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.5683832764625549, + "learning_rate": 2e-05, + "loss": 0.128, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.8846218585968018, + "learning_rate": 2e-05, + "loss": 0.2294, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.4786897897720337, + "learning_rate": 2e-05, + "loss": 0.6252, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.0091499090194702, + "learning_rate": 2e-05, + "loss": 0.2926, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.5462237596511841, + "learning_rate": 2e-05, + "loss": 0.4689, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.8900251388549805, + "learning_rate": 2e-05, + "loss": 0.3586, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.1280627250671387, + "learning_rate": 2e-05, + "loss": 0.2822, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.1584203243255615, + "learning_rate": 2e-05, + "loss": 0.2081, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.10125732421875, + "learning_rate": 2e-05, + "loss": 0.365, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.6134424209594727, + "learning_rate": 2e-05, + "loss": 0.2527, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5410845953622016.0, + "train_loss": 0.42002813339233397, + "train_runtime": 129.3934, + "train_samples_per_second": 3.091, + "train_steps_per_second": 0.773 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5410845953622016.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..e08ab30156db7de19043681789a6625305d3fce4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70343a7e54302748bdc8e092db9826cda1ecce709ae543ce39fbe042d70136b0 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..70d1d2d6fd528128cc0903839cff64d8608ed685 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e692ffba0e06842a0d55e58b2d4399b0308b920950f4cb5d09f625d1ed184328 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..ec6694b71e99e5757c13b5f3e74bacbd939a4379 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f906534287b0bddfd0af52aa9eab181914fa819f038eb0c7c36811f07f134cd9 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..b7152fb4f4ec1d900358f9ac70b65bd9e66ed2ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f732b90defc7422c6d154dbb4dff0dee6469630e3589b6665acbddd82ab9fa0d +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..f2d91ac23b272f5f08eca3dd9f077a0bb3ca4e03 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:159d45aa129260ea503bc35a73935421575fc8ee20fbc79e3348673690b3a742 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..3c5b09bfe815d954208278786c2df457ef88a7df --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c13df5d03b51d6eb964ad30171813e98459ccc35111498b13bfa68828ee19355 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..82427f50ff2469179a086924b664cc6a69db4bfb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f79132473eaff53474a14e14dbbe15281a905cd1361e215f96620570a3fdfad2 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..fa69e126ab5a71e4e9f8d51e7ba7576d4159ae78 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bd4aa86f1e72948ba95bd98252a9a1f9ca6ea767fac42bd43a3488f6b643d150 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..eef36bc37cf1654769e9f64753d0de1c86cfe8f0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/4_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.3453112840652466, + "learning_rate": 2e-05, + "loss": 0.1599, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.132049560546875, + "learning_rate": 2e-05, + "loss": 0.3725, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.7177079319953918, + "learning_rate": 2e-05, + "loss": 0.377, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.1375855207443237, + "learning_rate": 2e-05, + "loss": 0.3637, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.7680591940879822, + "learning_rate": 2e-05, + "loss": 0.224, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.4049601554870605, + "learning_rate": 2e-05, + "loss": 0.1168, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 1.795182466506958, + "learning_rate": 2e-05, + "loss": 0.3077, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 3.077700138092041, + "learning_rate": 2e-05, + "loss": 0.5691, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.2615911960601807, + "learning_rate": 2e-05, + "loss": 0.1729, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.5021533966064453, + "learning_rate": 2e-05, + "loss": 0.1832, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 1.9158926010131836, + "learning_rate": 2e-05, + "loss": 0.9795, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.8702447414398193, + "learning_rate": 2e-05, + "loss": 0.4986, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.3849501609802246, + "learning_rate": 2e-05, + "loss": 0.3267, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.3572231531143188, + "learning_rate": 2e-05, + "loss": 0.4395, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.1301411390304565, + "learning_rate": 2e-05, + "loss": 0.3774, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 1.4710966348648071, + "learning_rate": 2e-05, + "loss": 0.345, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.5909695625305176, + "learning_rate": 2e-05, + "loss": 0.1985, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 1.7595499753952026, + "learning_rate": 2e-05, + "loss": 0.4824, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.7000941038131714, + "learning_rate": 2e-05, + "loss": 0.2045, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.9647574424743652, + "learning_rate": 2e-05, + "loss": 0.6616, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.8519032597541809, + "learning_rate": 2e-05, + "loss": 0.4398, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.9542587399482727, + "learning_rate": 2e-05, + "loss": 0.2734, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.719585657119751, + "learning_rate": 2e-05, + "loss": 0.3818, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.8134914636611938, + "learning_rate": 2e-05, + "loss": 0.3435, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.9862194061279297, + "learning_rate": 2e-05, + "loss": 0.2208, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 1.1615169048309326, + "learning_rate": 2e-05, + "loss": 0.6509, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.904180645942688, + "learning_rate": 2e-05, + "loss": 0.2671, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.5389387607574463, + "learning_rate": 2e-05, + "loss": 0.1423, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.2976641654968262, + "learning_rate": 2e-05, + "loss": 0.3413, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 1.2798043489456177, + "learning_rate": 2e-05, + "loss": 0.4346, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.7560667991638184, + "learning_rate": 2e-05, + "loss": 0.2438, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.8956923484802246, + "learning_rate": 2e-05, + "loss": 0.4661, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.183799982070923, + "learning_rate": 2e-05, + "loss": 0.3533, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 1.172732949256897, + "learning_rate": 2e-05, + "loss": 0.3699, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.2292394638061523, + "learning_rate": 2e-05, + "loss": 0.4047, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.1442369222640991, + "learning_rate": 2e-05, + "loss": 0.2548, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5422228574752808, + "learning_rate": 2e-05, + "loss": 0.3146, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.2898380756378174, + "learning_rate": 2e-05, + "loss": 0.5573, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.5771578550338745, + "learning_rate": 2e-05, + "loss": 0.3517, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 2.2726993560791016, + "learning_rate": 2e-05, + "loss": 0.2866, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.7461665868759155, + "learning_rate": 2e-05, + "loss": 0.221, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.4335005283355713, + "learning_rate": 2e-05, + "loss": 0.3141, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.8735717535018921, + "learning_rate": 2e-05, + "loss": 0.1418, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.43066778779029846, + "learning_rate": 2e-05, + "loss": 0.2188, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.758823037147522, + "learning_rate": 2e-05, + "loss": 0.3583, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.6365742683410645, + "learning_rate": 2e-05, + "loss": 0.2551, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.1008527278900146, + "learning_rate": 2e-05, + "loss": 0.4255, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.8127040863037109, + "learning_rate": 2e-05, + "loss": 0.3618, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.619953989982605, + "learning_rate": 2e-05, + "loss": 0.0858, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 1.927673101425171, + "learning_rate": 2e-05, + "loss": 0.4988, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 6050325026832384.0, + "train_loss": 0.3467936706542969, + "train_runtime": 129.973, + "train_samples_per_second": 3.078, + "train_steps_per_second": 0.769 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 6050325026832384.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..5ab4c2a75b1ce879cff48f7ecd71e31fb8a952c2 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0aa8dbb8ada1c388a4d1f2160d644c9c4770a72515a14ae18cb19fcf7c4083dd +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..d85616c8dfcd1848be5ba37f6acbfac3f7e24937 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e95360434a6d70fcd6e3f0aa88423b833c81cc307cb33a344edb39ac2a320e35 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..43bd9b1827c3476f5bd4d52519dad4116454935e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5abb6747285e2c26e20278a03c35ef979ef152b4c7081f3d9581ed39a8b59aa7 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..24ca23fcffbdd10c704e7381a5c288d8cf3f8555 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8eb4f2040622e9c7b97a6bc63247c898ed5cf9dd3cacce4c3e5e2f53fc25829 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..a46e4d9a1805e7eb30fdb42696be61e5b9fcad8e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:22a39f32ae442c3da25a117a5498b87f49561462c3288db1c1d29410044fa0e1 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..5f1ed63f64121817a6414ca6062ef0584032ccef --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:19cca07727cc688860877c0d2e989b0cdf6bd33835df8d85a8e91ce07c0256b2 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..029d92ce3111cf25752083c0f91c2cfff51764a9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b3259a7e071370af432493e10e51095d1c25d08a9621c9781f43d4d8841a01d6 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..c71ce0c5bb3089f56df36037bac0f5bb2b049fc0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:de7efda77abbf34bbb3a1f9b4a37baeecaf4031baaceeca0538b3435398b685e +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..c33b64be12631317d175cbd2707b33f981d5d726 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/5_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.9637185335159302, + "learning_rate": 2e-05, + "loss": 0.1189, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 1.8156712055206299, + "learning_rate": 2e-05, + "loss": 0.2973, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.1615902185440063, + "learning_rate": 2e-05, + "loss": 0.1009, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.18979249894618988, + "learning_rate": 2e-05, + "loss": 0.0959, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.7038255929946899, + "learning_rate": 2e-05, + "loss": 0.122, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.1443534642457962, + "learning_rate": 2e-05, + "loss": 0.0486, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.005743575282394886, + "learning_rate": 2e-05, + "loss": 0.0359, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 0.07958142459392548, + "learning_rate": 2e-05, + "loss": 0.02, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.5263900756835938, + "learning_rate": 2e-05, + "loss": 0.0982, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.4334120750427246, + "learning_rate": 2e-05, + "loss": 0.1194, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.778258800506592, + "learning_rate": 2e-05, + "loss": 0.1253, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.1268651485443115, + "learning_rate": 2e-05, + "loss": 0.1391, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.5675013661384583, + "learning_rate": 2e-05, + "loss": 0.1794, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.5366666316986084, + "learning_rate": 2e-05, + "loss": 0.5917, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.09907013922929764, + "learning_rate": 2e-05, + "loss": 0.0348, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.7736491560935974, + "learning_rate": 2e-05, + "loss": 0.0499, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 1.6986557245254517, + "learning_rate": 2e-05, + "loss": 0.1054, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.20236420631408691, + "learning_rate": 2e-05, + "loss": 0.0331, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.8141397833824158, + "learning_rate": 2e-05, + "loss": 0.0374, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 2.192807674407959, + "learning_rate": 2e-05, + "loss": 0.159, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.16288095712661743, + "learning_rate": 2e-05, + "loss": 0.23, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 1.828477382659912, + "learning_rate": 2e-05, + "loss": 0.2355, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.4119545817375183, + "learning_rate": 2e-05, + "loss": 0.159, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.17235681414604187, + "learning_rate": 2e-05, + "loss": 0.8922, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 4.655712127685547, + "learning_rate": 2e-05, + "loss": 0.2608, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 2.3833858966827393, + "learning_rate": 2e-05, + "loss": 0.1453, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.1260771006345749, + "learning_rate": 2e-05, + "loss": 0.0114, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 1.6669312715530396, + "learning_rate": 2e-05, + "loss": 0.105, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.186197429895401, + "learning_rate": 2e-05, + "loss": 0.0162, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.7098574042320251, + "learning_rate": 2e-05, + "loss": 0.046, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.4444371461868286, + "learning_rate": 2e-05, + "loss": 0.0519, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.4043760299682617, + "learning_rate": 2e-05, + "loss": 0.6307, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 0.4836042523384094, + "learning_rate": 2e-05, + "loss": 0.0583, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.3592795431613922, + "learning_rate": 2e-05, + "loss": 0.0341, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 1.4283239841461182, + "learning_rate": 2e-05, + "loss": 0.2857, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 0.400084525346756, + "learning_rate": 2e-05, + "loss": 0.3132, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.08330664038658142, + "learning_rate": 2e-05, + "loss": 0.203, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.14160723984241486, + "learning_rate": 2e-05, + "loss": 0.009, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.09038477391004562, + "learning_rate": 2e-05, + "loss": 0.4529, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.14801684021949768, + "learning_rate": 2e-05, + "loss": 0.0669, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.44397568702697754, + "learning_rate": 2e-05, + "loss": 0.096, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.11191245913505554, + "learning_rate": 2e-05, + "loss": 0.0651, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.31379419565200806, + "learning_rate": 2e-05, + "loss": 0.0462, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.7242313027381897, + "learning_rate": 2e-05, + "loss": 0.0443, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.1290422677993774, + "learning_rate": 2e-05, + "loss": 0.0397, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 3.783607006072998, + "learning_rate": 2e-05, + "loss": 0.2218, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.4973742365837097, + "learning_rate": 2e-05, + "loss": 0.0265, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.1459741592407227, + "learning_rate": 2e-05, + "loss": 0.2214, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.05018605664372444, + "learning_rate": 2e-05, + "loss": 0.2559, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 5.330997467041016, + "learning_rate": 2e-05, + "loss": 1.1732, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5291433615425536.0, + "train_loss": 0.17817811965942382, + "train_runtime": 136.8533, + "train_samples_per_second": 2.923, + "train_steps_per_second": 0.731 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5291433615425536.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..44fee7b04c1a671f110277b7d90dbea392298c3b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:37a9f2187d2fb640db5892f3c06d2ee519a51b745f927ad33951f48328cb2550 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..531d54966e93dd280e3212e5d89824b1abb53c75 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6c3a547f1d83455b91a623c9edfcafe3b343a0dd75911583ee17e1059d80591a +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..e81ddc3ea882509dca414e187bded092db13fe79 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ee28bdb3e5279026f13d21a6301a350c7a109e05f5a7b57ff61a7d3bee09125c +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..475344b7144ed9c743217e0a62d466f01a776f62 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b4799d50978c47c4006c8ee3fda05c323df5acce1c33a9838c5503f77d8ff200 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..bf8c81b8f3f0ecb7822c4ff379c91318df1002aa --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d6c5e00496584f2c846ddd64b7bbeaf8e21f8c3a4f6699751c2635e2200a180 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..879aa733790c33ef2e05bf0986258a9d834736e3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7dd19a3d22f0e00ab387a7874ebe83bc1fedf797b2f6f95edd1c95509471eb3b +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..94b4853c4b6d0b4e173c166939a798bfe0d89cc3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6fd9a5dcd86825ac9b4c061c50d4b793bdef20e56ea1a10c67a353776fb6a7c7 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..969d77aa9d47584542a8026f6293e396d1ad667a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e036e29ecfa4087903ca64cb6fe614bae05567c22e61ff9c8394dbcb1ff191ea +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..6c2715ffb3a66034de8d5f5ac0bb21e46b9537d1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/6_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 3.064589262008667, + "learning_rate": 2e-05, + "loss": 0.3752, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 2.569004535675049, + "learning_rate": 2e-05, + "loss": 0.3982, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.7673529386520386, + "learning_rate": 2e-05, + "loss": 0.394, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 0.9663434624671936, + "learning_rate": 2e-05, + "loss": 0.5505, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 3.0199947357177734, + "learning_rate": 2e-05, + "loss": 0.52, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.6125103235244751, + "learning_rate": 2e-05, + "loss": 0.0558, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 2.0679080486297607, + "learning_rate": 2e-05, + "loss": 0.2921, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.9045493602752686, + "learning_rate": 2e-05, + "loss": 0.4228, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 2.740008592605591, + "learning_rate": 2e-05, + "loss": 0.6973, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 1.5434128046035767, + "learning_rate": 2e-05, + "loss": 0.412, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.277890682220459, + "learning_rate": 2e-05, + "loss": 0.39, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.4327271282672882, + "learning_rate": 2e-05, + "loss": 0.2969, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.6526048183441162, + "learning_rate": 2e-05, + "loss": 0.3651, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.3598532676696777, + "learning_rate": 2e-05, + "loss": 0.3755, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 3.399898052215576, + "learning_rate": 2e-05, + "loss": 0.9022, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 2.5105326175689697, + "learning_rate": 2e-05, + "loss": 0.371, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 2.0859789848327637, + "learning_rate": 2e-05, + "loss": 0.3644, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.9770312309265137, + "learning_rate": 2e-05, + "loss": 0.5312, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 1.5824873447418213, + "learning_rate": 2e-05, + "loss": 0.2852, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.1488003730773926, + "learning_rate": 2e-05, + "loss": 0.3595, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 1.6307090520858765, + "learning_rate": 2e-05, + "loss": 0.9417, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.7019992470741272, + "learning_rate": 2e-05, + "loss": 0.1998, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.9832942485809326, + "learning_rate": 2e-05, + "loss": 0.4437, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 1.5270015001296997, + "learning_rate": 2e-05, + "loss": 0.3009, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.8595147132873535, + "learning_rate": 2e-05, + "loss": 0.4537, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.8055856823921204, + "learning_rate": 2e-05, + "loss": 0.3318, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 1.5273380279541016, + "learning_rate": 2e-05, + "loss": 0.3598, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 3.0749666690826416, + "learning_rate": 2e-05, + "loss": 0.8135, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.0927863121032715, + "learning_rate": 2e-05, + "loss": 0.4734, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 2.371086597442627, + "learning_rate": 2e-05, + "loss": 0.586, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.02835750579834, + "learning_rate": 2e-05, + "loss": 0.3508, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.4880149364471436, + "learning_rate": 2e-05, + "loss": 0.4866, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.096529245376587, + "learning_rate": 2e-05, + "loss": 0.604, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 2.7746999263763428, + "learning_rate": 2e-05, + "loss": 0.4673, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 3.163508176803589, + "learning_rate": 2e-05, + "loss": 0.709, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.196310043334961, + "learning_rate": 2e-05, + "loss": 0.319, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 4.559057235717773, + "learning_rate": 2e-05, + "loss": 0.7285, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.912200927734375, + "learning_rate": 2e-05, + "loss": 0.5861, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.7412700653076172, + "learning_rate": 2e-05, + "loss": 0.3035, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 1.5687161684036255, + "learning_rate": 2e-05, + "loss": 0.2237, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 0.9691300392150879, + "learning_rate": 2e-05, + "loss": 0.3871, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.8654796481132507, + "learning_rate": 2e-05, + "loss": 0.6342, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.1085845232009888, + "learning_rate": 2e-05, + "loss": 0.2712, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.5613961219787598, + "learning_rate": 2e-05, + "loss": 0.502, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.862583041191101, + "learning_rate": 2e-05, + "loss": 0.3153, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 1.5156408548355103, + "learning_rate": 2e-05, + "loss": 0.9199, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 1.2614582777023315, + "learning_rate": 2e-05, + "loss": 0.3174, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.206857919692993, + "learning_rate": 2e-05, + "loss": 0.3384, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 2.434488534927368, + "learning_rate": 2e-05, + "loss": 0.5571, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 4.0444722175598145, + "learning_rate": 2e-05, + "loss": 1.0176, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 1.047142126845952e+16, + "train_loss": 0.46604217529296876, + "train_runtime": 186.9182, + "train_samples_per_second": 2.14, + "train_steps_per_second": 0.535 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 1.047142126845952e+16, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..c3cb6f4924f6026db5e93d6e928ded5fa0211e52 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dbebc1be29fa49e162820dbd4b278d79f4612ac4b2817d499e68298c8c6d671 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..96a8eb75f56e90cc78df35f4bc8cfdc45c1b273e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:87eb80622343e6f279b011744ad2abc0d44c62e36c3790903be09f134d32671d +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..57e1093cd52c5e4027adb3243158c9615490f507 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:030f62bf60b11a2c30032aefc9999e7df9050e8f0cb3bfbb44dc8e208f3dd898 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..c06ebd526408213d85c0fb237b226e683c0c33b1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:990abe62b229863d1a7e883f8ddc65c80b915d0731282c9931d58bbda2c52cb9 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..4d363dfcad181941147d0ae86b5a32f0d2a12ac0 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d791c932f2a8c515dcaa0cecc0cb10d2610d6e94755e88b1ef833792d629ef22 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..47f8fb8c244935a2cba32b039000dba86444e0e7 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b47df757edcb4248299aa0e9ee9206ae8dfc4fc7fad86e4746ba9035217cf38d +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..3f6c96a17f2383c7299aa42c1e07338692de7a9e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8f0ec2b0933a4537328dd1fac2e0c8fe2a56341e21d1461434f8d3c70b080f4c +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..3f8b60c31eca4b6d3df9258ebdde3dbe45db87f4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be92d9183c608f84e3c823c089f75a8063a4fbe8658eced322e15a6ac988b0e1 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..f6928886985d1c76f8764ab0d8ffac8694cbfc03 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/7_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 1.7729740142822266, + "learning_rate": 2e-05, + "loss": 0.1107, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 3.388115644454956, + "learning_rate": 2e-05, + "loss": 0.3786, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 1.9785577058792114, + "learning_rate": 2e-05, + "loss": 0.1904, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.6540454626083374, + "learning_rate": 2e-05, + "loss": 0.2095, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 1.9298994541168213, + "learning_rate": 2e-05, + "loss": 0.6282, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.12273844331502914, + "learning_rate": 2e-05, + "loss": 0.0403, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.38262739777565, + "learning_rate": 2e-05, + "loss": 0.0639, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 1.1514317989349365, + "learning_rate": 2e-05, + "loss": 0.3329, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.8205757141113281, + "learning_rate": 2e-05, + "loss": 0.1983, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.644108295440674, + "learning_rate": 2e-05, + "loss": 0.2732, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 0.11063311994075775, + "learning_rate": 2e-05, + "loss": 0.1941, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 3.6995646953582764, + "learning_rate": 2e-05, + "loss": 0.2978, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 0.5194157361984253, + "learning_rate": 2e-05, + "loss": 0.0325, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 2.1071629524230957, + "learning_rate": 2e-05, + "loss": 0.2634, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.43205615878105164, + "learning_rate": 2e-05, + "loss": 0.2947, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.6589717864990234, + "learning_rate": 2e-05, + "loss": 0.2767, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.41456761956214905, + "learning_rate": 2e-05, + "loss": 0.0308, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 2.4854793548583984, + "learning_rate": 2e-05, + "loss": 0.4358, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.6475552916526794, + "learning_rate": 2e-05, + "loss": 0.0569, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 3.861485719680786, + "learning_rate": 2e-05, + "loss": 0.6654, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 0.46356624364852905, + "learning_rate": 2e-05, + "loss": 0.4797, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 2.4162158966064453, + "learning_rate": 2e-05, + "loss": 0.2327, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 1.4475712776184082, + "learning_rate": 2e-05, + "loss": 0.322, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 0.16693030297756195, + "learning_rate": 2e-05, + "loss": 0.0085, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.3771916627883911, + "learning_rate": 2e-05, + "loss": 0.4371, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.723896861076355, + "learning_rate": 2e-05, + "loss": 0.4459, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.043569616973400116, + "learning_rate": 2e-05, + "loss": 0.3005, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.4647064507007599, + "learning_rate": 2e-05, + "loss": 0.0321, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 0.1501731425523758, + "learning_rate": 2e-05, + "loss": 0.1151, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 7.259007930755615, + "learning_rate": 2e-05, + "loss": 1.2305, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 0.5404440760612488, + "learning_rate": 2e-05, + "loss": 0.0703, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.6680632829666138, + "learning_rate": 2e-05, + "loss": 0.4868, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 1.4951984882354736, + "learning_rate": 2e-05, + "loss": 0.1459, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.5090456008911133, + "learning_rate": 2e-05, + "loss": 0.1313, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 2.162992000579834, + "learning_rate": 2e-05, + "loss": 0.3934, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 1.4258581399917603, + "learning_rate": 2e-05, + "loss": 0.3041, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 2.7312283515930176, + "learning_rate": 2e-05, + "loss": 0.4673, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.4119275212287903, + "learning_rate": 2e-05, + "loss": 0.3572, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 1.3404635190963745, + "learning_rate": 2e-05, + "loss": 0.3762, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.3218514919281006, + "learning_rate": 2e-05, + "loss": 0.78, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.2536154985427856, + "learning_rate": 2e-05, + "loss": 0.1497, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 1.2923953533172607, + "learning_rate": 2e-05, + "loss": 0.111, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 1.5413758754730225, + "learning_rate": 2e-05, + "loss": 0.3804, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 2.0926754474639893, + "learning_rate": 2e-05, + "loss": 0.3583, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 0.43409764766693115, + "learning_rate": 2e-05, + "loss": 0.2521, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 0.7976254820823669, + "learning_rate": 2e-05, + "loss": 0.1597, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.8634509444236755, + "learning_rate": 2e-05, + "loss": 0.1016, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.2198706716299057, + "learning_rate": 2e-05, + "loss": 0.0685, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 1.1303417682647705, + "learning_rate": 2e-05, + "loss": 0.1898, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 3.5226404666900635, + "learning_rate": 2e-05, + "loss": 0.3316, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5502714666549248.0, + "train_loss": 0.2838719272613525, + "train_runtime": 128.1285, + "train_samples_per_second": 3.122, + "train_steps_per_second": 0.78 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5502714666549248.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..645cc5791249c7b04589688c5e8245fcacac7196 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9c39b0b05cf90a8e9b8cfae5ac4e9f190a6b370cc9a5cbfd3596ed826e1232e4 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..dc641f5f341c095b2c3db2dd4a531056c1f29e0f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0b67d0153747c3b8b898adcbb6720dfddefb00e7685dd1f3a140ab6b7835ac51 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..04fc7b97e30d08c5f9a52084ba089d15b01e9fe4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63566df0038490bb31072c357a9b181e676aa4d37350705da4d0dc414905d49b +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..4d4c5b89ee972f4c97a9b312da6f945991873472 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7983c4d8126ad8fce6ecd03ab252c96c984e68315e5e030890f6159ebb4c09a8 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..d18c4941f4d0190db32c233abda6d5b7d7212fc8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:44d10a4310e88bab5ca98ff3f3981eaaf93560f5c9c1dda9291420ec4e91e74a +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..e602612cca6b45abf44ff3781b7065a594218022 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:42317b2a3721bf32f40cafa24a1874781a517f9a20d8ae4dca6a9e930f8bce62 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..dff3542c41cb87ec4b3fb7f9e3795fa02602fa30 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3964c14f3b8bc1c03e5dad63287833fb62ff7159bfaf82cd9d85e05adf2c834c +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..fede246b2158f0de9d8da09fa0e91c5230666293 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f2b5e081965bdc91554f6e58ff9c9c6b5a342da88dbc4e733958cc1371a4ede6 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..ce8d8fcd83830f32b35fba8613715cc748da2e88 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/8_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.05591817572712898, + "learning_rate": 2e-05, + "loss": 0.0331, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.05821004882454872, + "learning_rate": 2e-05, + "loss": 0.0794, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.4370066523551941, + "learning_rate": 2e-05, + "loss": 0.0509, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 2.820258140563965, + "learning_rate": 2e-05, + "loss": 0.2928, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.03945275768637657, + "learning_rate": 2e-05, + "loss": 0.0022, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 1.2746853828430176, + "learning_rate": 2e-05, + "loss": 0.0759, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.06721346825361252, + "learning_rate": 2e-05, + "loss": 0.1132, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 4.57942533493042, + "learning_rate": 2e-05, + "loss": 0.2443, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 1.2161158323287964, + "learning_rate": 2e-05, + "loss": 0.1112, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 3.6524734497070312, + "learning_rate": 2e-05, + "loss": 2.296, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.0981736183166504, + "learning_rate": 2e-05, + "loss": 0.4006, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 0.4628280997276306, + "learning_rate": 2e-05, + "loss": 0.4472, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 1.9978578090667725, + "learning_rate": 2e-05, + "loss": 0.1196, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 1.074185848236084, + "learning_rate": 2e-05, + "loss": 0.0613, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 1.5167150497436523, + "learning_rate": 2e-05, + "loss": 0.0576, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.056557487696409225, + "learning_rate": 2e-05, + "loss": 0.0118, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 3.0131587982177734, + "learning_rate": 2e-05, + "loss": 0.759, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 3.0006632804870605, + "learning_rate": 2e-05, + "loss": 0.3744, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 3.6181061267852783, + "learning_rate": 2e-05, + "loss": 0.6052, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.3154386281967163, + "learning_rate": 2e-05, + "loss": 0.5254, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 5.351212024688721, + "learning_rate": 2e-05, + "loss": 0.401, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 0.0025279216933995485, + "learning_rate": 2e-05, + "loss": 0.0677, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 2.9475443363189697, + "learning_rate": 2e-05, + "loss": 0.3987, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.0709035396575928, + "learning_rate": 2e-05, + "loss": 0.4412, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 0.25259989500045776, + "learning_rate": 2e-05, + "loss": 0.2133, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.04389140009880066, + "learning_rate": 2e-05, + "loss": 0.2017, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.378710001707077, + "learning_rate": 2e-05, + "loss": 0.2865, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.534584105014801, + "learning_rate": 2e-05, + "loss": 0.2747, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 1.1312546730041504, + "learning_rate": 2e-05, + "loss": 0.0746, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.07111427187919617, + "learning_rate": 2e-05, + "loss": 0.0081, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.2759335041046143, + "learning_rate": 2e-05, + "loss": 0.2206, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 3.2628672122955322, + "learning_rate": 2e-05, + "loss": 0.5791, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.0826921463012695, + "learning_rate": 2e-05, + "loss": 0.112, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.7393478751182556, + "learning_rate": 2e-05, + "loss": 0.0325, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.506921648979187, + "learning_rate": 2e-05, + "loss": 0.0242, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 2.7186739444732666, + "learning_rate": 2e-05, + "loss": 0.2126, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 1.156294584274292, + "learning_rate": 2e-05, + "loss": 0.0534, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 0.4439341127872467, + "learning_rate": 2e-05, + "loss": 0.0389, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.19657374918460846, + "learning_rate": 2e-05, + "loss": 0.1388, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 0.21871976554393768, + "learning_rate": 2e-05, + "loss": 0.0204, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 3.7503485679626465, + "learning_rate": 2e-05, + "loss": 0.4958, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 0.4864181876182556, + "learning_rate": 2e-05, + "loss": 0.024, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 0.1927212029695511, + "learning_rate": 2e-05, + "loss": 0.0106, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 0.6415959596633911, + "learning_rate": 2e-05, + "loss": 0.0446, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.3645117282867432, + "learning_rate": 2e-05, + "loss": 0.1523, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 4.426226615905762, + "learning_rate": 2e-05, + "loss": 0.5087, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 0.23949895799160004, + "learning_rate": 2e-05, + "loss": 0.0087, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 2.158336877822876, + "learning_rate": 2e-05, + "loss": 0.1513, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.20911116898059845, + "learning_rate": 2e-05, + "loss": 0.0108, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.14650678634643555, + "learning_rate": 2e-05, + "loss": 0.0092, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5329528591220736.0, + "train_loss": 0.23754099369049073, + "train_runtime": 126.8303, + "train_samples_per_second": 3.154, + "train_steps_per_second": 0.788 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5329528591220736.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth new file mode 100644 index 0000000000000000000000000000000000000000..6aff84f3442c18914d92ba02d19ec1205d4de176 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round10.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2cdce81d7ccdf4c6cbb01bc4398c42a9a4156cedd7d5a3be00e71071cc8281ea +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth new file mode 100644 index 0000000000000000000000000000000000000000..3f575478bfbcf31d649894ca8bb54fbc6e1f5bce --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round12.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3514a477c4a7c48ef5c158eeb723025fd937926892ab6f6776225bdfa852ee8 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth new file mode 100644 index 0000000000000000000000000000000000000000..0655568efa49683fba37bb8a27c9ae3d8691647b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round15.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:148f20bbca7d952d496c931220d8d65a5127e60585de722f9bbebaf59d24db3a +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth new file mode 100644 index 0000000000000000000000000000000000000000..6c28f48d9f02641897b6ecd6c3ef19abdd5b4cc4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round17.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:938609c889f894c9db39101892fc6866a93e92473373557edd87720763fb8fe8 +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth new file mode 100644 index 0000000000000000000000000000000000000000..98b9f9aa558bf6e1bc1feb5cc759121f4d931b67 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round2.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9b44052c1abf77090db7a41040c3e38f85b1b96fb20a3647a10dadae03a80e58 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth new file mode 100644 index 0000000000000000000000000000000000000000..e8edb860ae4bb133fbefefc8750c157b1df23e2f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round20.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:439e5c3381333186ebae66599fb27f339358e0df7891bfbb3ce00c663a14e67e +size 794708086 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth new file mode 100644 index 0000000000000000000000000000000000000000..b262946d565b4b7b8756614fb274f01fd6ce91a5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round5.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9a891dbc7c3a20e8b122a0ae9e6b55ffe91ed1e87a3b085be64e9f878dcb5808 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth new file mode 100644 index 0000000000000000000000000000000000000000..ccee300cbe2a147dce49edd042b3e79c84c0a4dd --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_client_model_round7.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b686bc696b1bc49d0eb3db112997b55243d15c082d1222d1f1f9816430bc9ad3 +size 794706058 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json new file mode 100644 index 0000000000000000000000000000000000000000..3e1bb88c5b828408584eff29971d39bfee772940 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/9_trainer_state.json @@ -0,0 +1,392 @@ +{ + "best_metric": null, + "best_model_checkpoint": null, + "epoch": 1.0, + "eval_steps": 500, + "global_step": 100, + "is_hyper_param_search": false, + "is_local_process_zero": true, + "is_world_process_zero": true, + "log_history": [ + { + "epoch": 0.02, + "grad_norm": 0.23775307834148407, + "learning_rate": 2e-05, + "loss": 0.3294, + "step": 2 + }, + { + "epoch": 0.04, + "grad_norm": 0.47008663415908813, + "learning_rate": 2e-05, + "loss": 0.6272, + "step": 4 + }, + { + "epoch": 0.06, + "grad_norm": 0.6493474841117859, + "learning_rate": 2e-05, + "loss": 0.0642, + "step": 6 + }, + { + "epoch": 0.08, + "grad_norm": 1.224001169204712, + "learning_rate": 2e-05, + "loss": 0.1311, + "step": 8 + }, + { + "epoch": 0.1, + "grad_norm": 0.18383054435253143, + "learning_rate": 2e-05, + "loss": 0.3403, + "step": 10 + }, + { + "epoch": 0.12, + "grad_norm": 0.5218638181686401, + "learning_rate": 2e-05, + "loss": 0.1457, + "step": 12 + }, + { + "epoch": 0.14, + "grad_norm": 0.057604387402534485, + "learning_rate": 2e-05, + "loss": 0.39, + "step": 14 + }, + { + "epoch": 0.16, + "grad_norm": 2.4070374965667725, + "learning_rate": 2e-05, + "loss": 0.5833, + "step": 16 + }, + { + "epoch": 0.18, + "grad_norm": 0.46707817912101746, + "learning_rate": 2e-05, + "loss": 0.1485, + "step": 18 + }, + { + "epoch": 0.2, + "grad_norm": 2.2831714153289795, + "learning_rate": 2e-05, + "loss": 0.2023, + "step": 20 + }, + { + "epoch": 0.22, + "grad_norm": 2.0630602836608887, + "learning_rate": 2e-05, + "loss": 0.7499, + "step": 22 + }, + { + "epoch": 0.24, + "grad_norm": 2.814023017883301, + "learning_rate": 2e-05, + "loss": 0.4632, + "step": 24 + }, + { + "epoch": 0.26, + "grad_norm": 2.094834804534912, + "learning_rate": 2e-05, + "loss": 0.4937, + "step": 26 + }, + { + "epoch": 0.28, + "grad_norm": 0.9579790234565735, + "learning_rate": 2e-05, + "loss": 0.0741, + "step": 28 + }, + { + "epoch": 0.3, + "grad_norm": 0.17668387293815613, + "learning_rate": 2e-05, + "loss": 0.1841, + "step": 30 + }, + { + "epoch": 0.32, + "grad_norm": 0.11519224941730499, + "learning_rate": 2e-05, + "loss": 0.0068, + "step": 32 + }, + { + "epoch": 0.34, + "grad_norm": 0.6470044851303101, + "learning_rate": 2e-05, + "loss": 0.5621, + "step": 34 + }, + { + "epoch": 0.36, + "grad_norm": 0.8433932065963745, + "learning_rate": 2e-05, + "loss": 0.1763, + "step": 36 + }, + { + "epoch": 0.38, + "grad_norm": 0.2445361614227295, + "learning_rate": 2e-05, + "loss": 0.0623, + "step": 38 + }, + { + "epoch": 0.4, + "grad_norm": 1.7083752155303955, + "learning_rate": 2e-05, + "loss": 0.3568, + "step": 40 + }, + { + "epoch": 0.42, + "grad_norm": 2.2258684635162354, + "learning_rate": 2e-05, + "loss": 0.2952, + "step": 42 + }, + { + "epoch": 0.44, + "grad_norm": 3.2891528606414795, + "learning_rate": 2e-05, + "loss": 0.6074, + "step": 44 + }, + { + "epoch": 0.46, + "grad_norm": 0.6204328536987305, + "learning_rate": 2e-05, + "loss": 0.0576, + "step": 46 + }, + { + "epoch": 0.48, + "grad_norm": 3.2032012939453125, + "learning_rate": 2e-05, + "loss": 0.4326, + "step": 48 + }, + { + "epoch": 0.5, + "grad_norm": 1.139670491218567, + "learning_rate": 2e-05, + "loss": 0.056, + "step": 50 + }, + { + "epoch": 0.52, + "grad_norm": 0.5201926231384277, + "learning_rate": 2e-05, + "loss": 0.0897, + "step": 52 + }, + { + "epoch": 0.54, + "grad_norm": 0.5032634139060974, + "learning_rate": 2e-05, + "loss": 0.4951, + "step": 54 + }, + { + "epoch": 0.56, + "grad_norm": 0.1492750197649002, + "learning_rate": 2e-05, + "loss": 0.0121, + "step": 56 + }, + { + "epoch": 0.58, + "grad_norm": 2.0449228286743164, + "learning_rate": 2e-05, + "loss": 0.3633, + "step": 58 + }, + { + "epoch": 0.6, + "grad_norm": 0.5991446375846863, + "learning_rate": 2e-05, + "loss": 0.0805, + "step": 60 + }, + { + "epoch": 0.62, + "grad_norm": 2.067772150039673, + "learning_rate": 2e-05, + "loss": 0.1803, + "step": 62 + }, + { + "epoch": 0.64, + "grad_norm": 1.2915898561477661, + "learning_rate": 2e-05, + "loss": 0.1985, + "step": 64 + }, + { + "epoch": 0.66, + "grad_norm": 2.1832423210144043, + "learning_rate": 2e-05, + "loss": 0.18, + "step": 66 + }, + { + "epoch": 0.68, + "grad_norm": 0.802361011505127, + "learning_rate": 2e-05, + "loss": 0.0829, + "step": 68 + }, + { + "epoch": 0.7, + "grad_norm": 0.14178869128227234, + "learning_rate": 2e-05, + "loss": 0.023, + "step": 70 + }, + { + "epoch": 0.72, + "grad_norm": 4.70101261138916, + "learning_rate": 2e-05, + "loss": 0.4852, + "step": 72 + }, + { + "epoch": 0.74, + "grad_norm": 0.5570341348648071, + "learning_rate": 2e-05, + "loss": 0.1658, + "step": 74 + }, + { + "epoch": 0.76, + "grad_norm": 2.144829750061035, + "learning_rate": 2e-05, + "loss": 0.3092, + "step": 76 + }, + { + "epoch": 0.78, + "grad_norm": 0.39063170552253723, + "learning_rate": 2e-05, + "loss": 0.1239, + "step": 78 + }, + { + "epoch": 0.8, + "grad_norm": 3.634079933166504, + "learning_rate": 2e-05, + "loss": 0.4109, + "step": 80 + }, + { + "epoch": 0.82, + "grad_norm": 1.8974275588989258, + "learning_rate": 2e-05, + "loss": 0.2895, + "step": 82 + }, + { + "epoch": 0.84, + "grad_norm": 2.1156625747680664, + "learning_rate": 2e-05, + "loss": 0.2263, + "step": 84 + }, + { + "epoch": 0.86, + "grad_norm": 2.2234652042388916, + "learning_rate": 2e-05, + "loss": 0.2163, + "step": 86 + }, + { + "epoch": 0.88, + "grad_norm": 1.485494613647461, + "learning_rate": 2e-05, + "loss": 0.3678, + "step": 88 + }, + { + "epoch": 0.9, + "grad_norm": 1.015352725982666, + "learning_rate": 2e-05, + "loss": 0.2766, + "step": 90 + }, + { + "epoch": 0.92, + "grad_norm": 2.3957250118255615, + "learning_rate": 2e-05, + "loss": 0.2656, + "step": 92 + }, + { + "epoch": 0.94, + "grad_norm": 2.6774632930755615, + "learning_rate": 2e-05, + "loss": 0.4917, + "step": 94 + }, + { + "epoch": 0.96, + "grad_norm": 0.3081165552139282, + "learning_rate": 2e-05, + "loss": 0.0526, + "step": 96 + }, + { + "epoch": 0.98, + "grad_norm": 0.8693525791168213, + "learning_rate": 2e-05, + "loss": 0.2766, + "step": 98 + }, + { + "epoch": 1.0, + "grad_norm": 0.16221743822097778, + "learning_rate": 2e-05, + "loss": 0.0688, + "step": 100 + }, + { + "epoch": 1.0, + "step": 100, + "total_flos": 5293891360129024.0, + "train_loss": 0.265445671081543, + "train_runtime": 127.4363, + "train_samples_per_second": 3.139, + "train_steps_per_second": 0.785 + } + ], + "logging_steps": 2, + "max_steps": 100, + "num_input_tokens_seen": 0, + "num_train_epochs": 1, + "save_steps": 500, + "stateful_callbacks": { + "TrainerControl": { + "args": { + "should_epoch_stop": false, + "should_evaluate": false, + "should_log": false, + "should_save": false, + "should_training_stop": false + }, + "attributes": {} + } + }, + "total_flos": 5293891360129024.0, + "train_batch_size": 2, + "trial_name": null, + "trial_params": null +} diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2efc0066480f002b8cceece896f7055865c8ae8f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:84e0d0d14362b8a53540c149bfdf3ad87bd81fefc2437a5a8fc577f8bbe796bc +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..6fb0c5e5355623c152966566c94dae6340813595 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5c775bb55ca1f83f38a0018b68ff84cb123f4cfc6bb67d2c425af418aff5d090 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_0/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a378e644f3a3028d2a57b81be98fe4327dbd6285 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c65df32b7c12a99f062509d972c24760fa82083161342273699f00eed5718c72 +size 841733872 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..48bd1cb395e7d769362f0b776e91087df416a409 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b1f085e020a2c73618a3380ca6499c3311893381015bda136ac41d17c27a0847 +size 895436332 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_1/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ea044582952242ffac4806f1625bd551114b7c61 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cd6aa6769f4e8739d59db2fc25078802d38314f4f28de0d92dc8542cb6c16ec7 +size 841733872 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..dc65420d5b4625b810780991830173df5f0bef8f --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9dbb188f9365855b5004149846ad2a1a9a1aa2a7a601ecdf1865957d4c3fda4f +size 895436332 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_10/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..462ab9a6bed7e79231f119caca51f605f978352a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be7b4da93d4380edd8f9ec1970e4d1ad28ef610e29f4792c33b2a8cd96b8cf9c +size 841733872 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..5cc971abdca2fc966bdbd29b22393757fc4650bc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:33547fe2b036817eeaacbd58677b5a2d8b01dce04b8b1783f8893fbf2376e743 +size 895436332 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_11/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..108a1b45efcd39f0e85a99589008230719d72d09 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a51d54ae684d30fc51fda845ba908658bb042737caf33cb5e23850901f4b1b03 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..5b162a0988f6c33180983628b55fd5b19977787e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f5db4b26b18a2fe9066e3ebc60e0daeb0bc87dc13a63ab5c0f2d59a80ca29928 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_12/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..03d73e693b865f58c7586ed9c7612022ac7269d4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:91ed5297edb0260e6e76e3ff2de5ce373d5f1341ff42262de4e48915d2041585 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..40c6ba9c2f9f2f9743657995ad929db7d519010a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b21b815acb82416e71a45a2f635e37c85bcf9079aa5871fec2dd012cc35244ce +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_13/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..264cd5d238e97c0defc34cd8bc06c6ef81015ffe --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:71702cc4adef3fb1eec06008f2bcdd37baf94380f63132bf9139e55cc869fa93 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..52d5ff72cc28de20c85a6eb95b4a7dca6034de52 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ec137bc02a806e640f38f935b437cacc7e2a609534a90f134acb3f3f2991d5fa +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_14/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..23c974e60a21db550fdf49ca683084c7b2d62fdb --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:59e27bc0e1e3eeffd088afc5c8f86196e33ebe490a2326d47578e1efb368682e +size 841733872 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..7a116ff2e948f62fdddc81ce081e16363a33830c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a19776784936287ba21647e87e90a79d6e2fc0071a80b8c26bab7d14a04ba06a +size 895436332 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_15/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..528028809695cc69981c551d5b17fc75cc674dc8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a9bc1a3cb0bb282d04599a3eb2068c26243ea26fb4fb4bec18c435588e55f6e4 +size 841733872 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9f9907df333bff8ff357efb207a602959df4c393 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3ab6f6499c9521b5f8774f697ad9f21efb3a9a76664c466e2e36c685389762aa +size 895436332 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_16/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..f02c5880fe3c2bb66179a5643715ff94955b0b09 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c2101c40f14c9c91c051232879afa55874de095484d7b131492120a631769d72 +size 841733872 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..40300c5fc1b520056d83db44847a9efb760a0da1 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b676bdc7099677685a26f406232ee914bba4e89126ff0230fea25bf5851f0efa +size 895436332 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_17/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..2d2b072b03f9707dd243a14aa1ab266b6bf6b809 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eef537250b91176871d0bb2ddb338d7d21689b05af34f9068b82d0f9d4f7b87f +size 841733872 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..48785ae4a480b84f6be634f5a6c5db17a6000430 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bec3c7bcddf925627c1bad59ad76a5ce0a70dc4f9c8cfa230f08cdf07af14a60 +size 895436332 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_18/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..b19c2a74722fcbd745b047babb33a60816e93398 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9045f3d3eabc12730931582d4a8763f18f1186e7551f561695eed7865553e5f7 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..a7ab9ca3545899ca37f38cbd864885919e02bc08 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:49bdcc6283352cc5183411cbea0bd520c1c9bf1aabcf820a3ab46103768c8987 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_19/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4523552bfa7b1b948524046ba47840914c7e2680 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09c9b01ce47fc8cd98845d9232ff17bfa526825dd6db62e718a9819016fd0c17 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0c2fd7fab7a1b611522c01c4f3e583799050fd5c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6f29cccaae4a31b25bc9724273f239f44df026e64a9d9d5878cf984c2cbfc4ca +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_2/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..4d2b9fa0e926ce58287ee27eb10ebe331d411e38 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:07638534388a3afc4e2a9320e7b53359bf0512b932bcfcfe054204f1031d8868 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..dd22d3258d3cafa15d9bbff1ddc694bd61459467 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b11f561ea46ebf9433c8d9c874234fd81288893c391d46f86fff09d436cc050f +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_20/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c61af160a0e6c243bad39d31fae69f6721021b4b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:54d47e41ad4ce9dd7ab4758ca40a1dfb2132a8990599500248e13f028a8ccd2e +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1aa4781a4c2fe77197ada34adc41db6a08c40fa4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a501d2e463916d94d631edc673408dfcb14ddb04c8b9b45722e64ca1d531795d +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_3/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..0da092915ab0dd9303649241ba1369e7cd3773a3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8dde4a19a53bfdfaf57e967351091c4c3fcc339def20606d521443af3ebba7d4 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e79256073eb7d04895f98fadc9a5fd40dfa95b85 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a7b3ab7628e01ace4f621640277ad03529b11a47c458cafc5fa50fa45c23f501 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_4/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..ef78b82134f9657149819829ccbdedecc08a7acc --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3cea217e9ff67a22e633d8107f79a544886f06fac474577640c7174e9ee998b8 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..e4a398e3b8d197a9e25290a9df64d932ded74c43 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6e3ad28f960ad178cb5e3a545db65fa1bd6cf75b66b73fc116811d66f7e89fd5 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_5/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..240441e31f42bdea26955a8c12e6f09997f31245 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:113f63ca515451a39955e6a782225b95ac598b14259644d6b97bafe65d98cbce +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..9f83c19eee8100bf831b75e60b2fc821488349af --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3a34505c90f5cd8887fe269c0f3988a00838363bffbf7f06e31300902e7a3d81 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_6/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..1b992062613eb2fd077fbf43e0504baa41b05011 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4ac5713acc6d00e247a8a20159321a389057f6809b65d3750c0887688fbea1bd +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d8de72f781b37ea25f54b714d5447f782d921989 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:69328fec83bd88e99306462df9b28871d32cb22a46445921c72fc11a952754fa +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_7/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..d942f45d8b3f455392f2047f0fd99015ec376a2b --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:04fa5ce8fc4cbcab06f80be46d63e0615c77657b3bf97d6c0eeceb710afe2510 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..54244099e2311d045038ea58813581d407e4e491 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b76f7b223afbf0245b46478f2ce9cbcdd2c9f3695b697c6d743f49cdec03c491 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_8/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..c86dd11afc45ac065fdac3a19e080e31cb65ef6d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/bf16_zero_pp_rank_0_mp_rank_00_optim_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b6ed3e54d5b7fc3a4133f6f66d3b4c516282111da15261b93bb4146669f6d222 +size 2057551984 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt new file mode 100644 index 0000000000000000000000000000000000000000..752aa88bb926229654c31b5ea0f9994c5ff47e6a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/global_step100/mp_rank_00_model_states.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d51e2777068573f4f42b11a9a4274beb40e3126d5c192be25824ce88f79e2f54 +size 1131744388 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/latest b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/latest new file mode 100644 index 0000000000000000000000000000000000000000..744ae7dbad571b6f37ec6c7066549494261bb59e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/latest @@ -0,0 +1 @@ +global_step100 \ No newline at end of file diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt new file mode 100644 index 0000000000000000000000000000000000000000..1dca79e108964eb087b0927c948e86e5b07cc7ec --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/scheduler.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:93bd89eea21ce12821280c3bc8a42ce07a0c32c7007fdd71c0f147ae8ac79e41 +size 1064 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py new file mode 100755 index 0000000000000000000000000000000000000000..24cc342e78d1a006c782b3a4cd68d9ce786d8fd8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/client_9/zero_to_fp32.py @@ -0,0 +1,604 @@ +#!/usr/bin/env python + +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +# This script extracts fp32 consolidated weights from a zero 1, 2 and 3 DeepSpeed checkpoints. It gets +# copied into the top level checkpoint dir, so the user can easily do the conversion at any point in +# the future. Once extracted, the weights don't require DeepSpeed and can be used in any +# application. +# +# example: python zero_to_fp32.py . pytorch_model.bin + +import argparse +import torch +import glob +import math +import os +import re +from collections import OrderedDict +from dataclasses import dataclass + +# while this script doesn't use deepspeed to recover data, since the checkpoints are pickled with +# DeepSpeed data structures it has to be available in the current python environment. +from deepspeed.utils import logger +from deepspeed.checkpoint.constants import (DS_VERSION, OPTIMIZER_STATE_DICT, SINGLE_PARTITION_OF_FP32_GROUPS, + FP32_FLAT_GROUPS, ZERO_STAGE, PARTITION_COUNT, PARAM_SHAPES, BUFFER_NAMES, + FROZEN_PARAM_SHAPES, FROZEN_PARAM_FRAGMENTS) + + +@dataclass +class zero_model_state: + buffers: dict() + param_shapes: dict() + shared_params: list + ds_version: int + frozen_param_shapes: dict() + frozen_param_fragments: dict() + + +debug = 0 + +# load to cpu +device = torch.device('cpu') + + +def atoi(text): + return int(text) if text.isdigit() else text + + +def natural_keys(text): + ''' + alist.sort(key=natural_keys) sorts in human order + http://nedbatchelder.com/blog/200712/human_sorting.html + (See Toothy's implementation in the comments) + ''' + return [atoi(c) for c in re.split(r'(\d+)', text)] + + +def get_model_state_file(checkpoint_dir, zero_stage): + if not os.path.isdir(checkpoint_dir): + raise FileNotFoundError(f"Directory '{checkpoint_dir}' doesn't exist") + + # there should be only one file + if zero_stage <= 2: + file = os.path.join(checkpoint_dir, "mp_rank_00_model_states.pt") + elif zero_stage == 3: + file = os.path.join(checkpoint_dir, "zero_pp_rank_0_mp_rank_00_model_states.pt") + + if not os.path.exists(file): + raise FileNotFoundError(f"can't find model states file at '{file}'") + + return file + + +def get_checkpoint_files(checkpoint_dir, glob_pattern): + # XXX: need to test that this simple glob rule works for multi-node setup too + ckpt_files = sorted(glob.glob(os.path.join(checkpoint_dir, glob_pattern)), key=natural_keys) + + if len(ckpt_files) == 0: + raise FileNotFoundError(f"can't find {glob_pattern} files in directory '{checkpoint_dir}'") + + return ckpt_files + + +def get_optim_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_optim_states.pt") + + +def get_model_state_files(checkpoint_dir): + return get_checkpoint_files(checkpoint_dir, "*_model_states.pt") + + +def parse_model_states(files): + zero_model_states = [] + for file in files: + state_dict = torch.load(file, map_location=device) + + if BUFFER_NAMES not in state_dict: + raise ValueError(f"{file} is not a model state checkpoint") + buffer_names = state_dict[BUFFER_NAMES] + if debug: + print("Found buffers:", buffer_names) + + # recover just the buffers while restoring them to fp32 if they were saved in fp16 + buffers = {k: v.float() for k, v in state_dict["module"].items() if k in buffer_names} + param_shapes = state_dict[PARAM_SHAPES] + + # collect parameters that are included in param_shapes + param_names = [] + for s in param_shapes: + for name in s.keys(): + param_names.append(name) + + # update with frozen parameters + frozen_param_shapes = state_dict.get(FROZEN_PARAM_SHAPES, None) + if frozen_param_shapes is not None: + if debug: + print(f"Found frozen_param_shapes: {frozen_param_shapes}") + param_names += list(frozen_param_shapes.keys()) + + # handle shared params + shared_params = [[k, v] for k, v in state_dict["shared_params"].items()] + + ds_version = state_dict.get(DS_VERSION, None) + + frozen_param_fragments = state_dict.get(FROZEN_PARAM_FRAGMENTS, None) + + z_model_state = zero_model_state(buffers=buffers, + param_shapes=param_shapes, + shared_params=shared_params, + ds_version=ds_version, + frozen_param_shapes=frozen_param_shapes, + frozen_param_fragments=frozen_param_fragments) + zero_model_states.append(z_model_state) + + return zero_model_states + + +def parse_optim_states(files, ds_checkpoint_dir): + + total_files = len(files) + state_dicts = [] + for f in files: + state_dict = torch.load(f, map_location=device) + # immediately discard the potentially huge 2 optimizer states as we only care for fp32 master weights + # and also handle the case where it was already removed by another helper script + state_dict["optimizer_state_dict"].pop("optimizer_state_dict", None) + state_dicts.append(state_dict) + + if not ZERO_STAGE in state_dicts[0][OPTIMIZER_STATE_DICT]: + raise ValueError(f"{files[0]} is not a zero checkpoint") + zero_stage = state_dicts[0][OPTIMIZER_STATE_DICT][ZERO_STAGE] + world_size = state_dicts[0][OPTIMIZER_STATE_DICT][PARTITION_COUNT] + + # For ZeRO-2 each param group can have different partition_count as data parallelism for expert + # parameters can be different from data parallelism for non-expert parameters. So we can just + # use the max of the partition_count to get the dp world_size. + + if type(world_size) is list: + world_size = max(world_size) + + if world_size != total_files: + raise ValueError( + f"Expected {world_size} of '*_optim_states.pt' under '{ds_checkpoint_dir}' but found {total_files} files. " + "Possibly due to an overwrite of an old checkpoint, or a checkpoint didn't get saved by one or more processes." + ) + + # the groups are named differently in each stage + if zero_stage <= 2: + fp32_groups_key = SINGLE_PARTITION_OF_FP32_GROUPS + elif zero_stage == 3: + fp32_groups_key = FP32_FLAT_GROUPS + else: + raise ValueError(f"unknown zero stage {zero_stage}") + + if zero_stage <= 2: + fp32_flat_groups = [state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key] for i in range(len(state_dicts))] + elif zero_stage == 3: + # if there is more than one param group, there will be multiple flattened tensors - one + # flattened tensor per group - for simplicity merge them into a single tensor + # + # XXX: could make the script more memory efficient for when there are multiple groups - it + # will require matching the sub-lists of param_shapes for each param group flattened tensor + + fp32_flat_groups = [ + torch.cat(state_dicts[i][OPTIMIZER_STATE_DICT][fp32_groups_key], 0) for i in range(len(state_dicts)) + ] + + return zero_stage, world_size, fp32_flat_groups + + +def _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters): + """ + Returns fp32 state_dict reconstructed from ds checkpoint + + Args: + - ``ds_checkpoint_dir``: path to the deepspeed checkpoint folder (where the optimizer files are) + + """ + print(f"Processing zero checkpoint '{ds_checkpoint_dir}'") + + optim_files = get_optim_files(ds_checkpoint_dir) + zero_stage, world_size, fp32_flat_groups = parse_optim_states(optim_files, ds_checkpoint_dir) + print(f"Detected checkpoint of type zero stage {zero_stage}, world_size: {world_size}") + + model_files = get_model_state_files(ds_checkpoint_dir) + + zero_model_states = parse_model_states(model_files) + print(f'Parsing checkpoint created by deepspeed=={zero_model_states[0].ds_version}') + + if zero_stage <= 2: + return _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + elif zero_stage == 3: + return _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters) + + +def _zero2_merge_frozen_params(state_dict, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + frozen_param_fragments = zero_model_states[0].frozen_param_fragments + + if debug: + num_elem = sum(s.numel() for s in frozen_param_shapes.values()) + print(f'rank 0: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in frozen_param_fragments.values()]) + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + state_dict[name] = frozen_param_fragments[name] + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _has_callable(obj, fn): + attr = getattr(obj, fn, None) + return callable(attr) + + +def _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + + # Reconstruction protocol: + # + # XXX: document this + + if debug: + for i in range(world_size): + for j in range(len(fp32_flat_groups[0])): + print(f"{FP32_FLAT_GROUPS}[{i}][{j}].shape={fp32_flat_groups[i][j].shape}") + + # XXX: memory usage doubles here (zero2) + num_param_groups = len(fp32_flat_groups[0]) + merged_single_partition_of_fp32_groups = [] + for i in range(num_param_groups): + merged_partitions = [sd[i] for sd in fp32_flat_groups] + full_single_fp32_vector = torch.cat(merged_partitions, 0) + merged_single_partition_of_fp32_groups.append(full_single_fp32_vector) + avail_numel = sum( + [full_single_fp32_vector.numel() for full_single_fp32_vector in merged_single_partition_of_fp32_groups]) + + if debug: + wanted_params = sum([len(shapes) for shapes in param_shapes]) + wanted_numel = sum([sum(shape.numel() for shape in shapes.values()) for shapes in param_shapes]) + # not asserting if there is a mismatch due to possible padding + print(f"Have {avail_numel} numels to process.") + print(f"Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + total_numel = 0 + total_params = 0 + for shapes, full_single_fp32_vector in zip(param_shapes, merged_single_partition_of_fp32_groups): + offset = 0 + avail_numel = full_single_fp32_vector.numel() + for name, shape in shapes.items(): + + unpartitioned_numel = shape.numel() if _has_callable(shape, 'numel') else math.prod(shape) + total_numel += unpartitioned_numel + total_params += 1 + + if debug: + print(f"{name} full shape: {shape} unpartitioned numel {unpartitioned_numel} ") + state_dict[name] = full_single_fp32_vector.narrow(0, offset, unpartitioned_numel).view(shape) + offset += unpartitioned_numel + + # Z2 started to align to 2*world_size to improve nccl performance. Therefore both offset and + # avail_numel can differ by anywhere between 0..2*world_size. Due to two unrelated complex + # paddings performed in the code it's almost impossible to predict the exact numbers w/o the + # live optimizer object, so we are checking that the numbers are within the right range + align_to = 2 * world_size + + def zero2_align(x): + return align_to * math.ceil(x / align_to) + + if debug: + print(f"original offset={offset}, avail_numel={avail_numel}") + + offset = zero2_align(offset) + avail_numel = zero2_align(avail_numel) + + if debug: + print(f"aligned offset={offset}, avail_numel={avail_numel}") + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero2_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero2_merge_frozen_params(state_dict, zero_model_states) + + _zero2_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def zero3_partitioned_param_info(unpartitioned_numel, world_size): + remainder = unpartitioned_numel % world_size + padding_numel = (world_size - remainder) if remainder else 0 + partitioned_numel = math.ceil(unpartitioned_numel / world_size) + return partitioned_numel, padding_numel + + +def _zero3_merge_frozen_params(state_dict, world_size, zero_model_states): + if zero_model_states[0].frozen_param_shapes is None or len(zero_model_states[0].frozen_param_shapes) == 0: + return + + if debug: + for i in range(world_size): + num_elem = sum(s.numel() for s in zero_model_states[i].frozen_param_fragments.values()) + print(f'rank {i}: {FROZEN_PARAM_SHAPES}.numel = {num_elem}') + + frozen_param_shapes = zero_model_states[0].frozen_param_shapes + wanted_params = len(frozen_param_shapes) + wanted_numel = sum(s.numel() for s in frozen_param_shapes.values()) + avail_numel = sum([p.numel() for p in zero_model_states[0].frozen_param_fragments.values()]) * world_size + print(f'Frozen params: Have {avail_numel} numels to process.') + print(f'Frozen params: Need {wanted_numel} numels in {wanted_params} params') + + total_params = 0 + total_numel = 0 + for name, shape in zero_model_states[0].frozen_param_shapes.items(): + total_params += 1 + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + + param_frags = tuple(model_state.frozen_param_fragments[name] for model_state in zero_model_states) + state_dict[name] = torch.cat(param_frags, 0).narrow(0, 0, unpartitioned_numel).view(shape) + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Frozen params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + print(f"Reconstructed Frozen fp32 state dict with {total_params} params {total_numel} elements") + + +def _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states): + param_shapes = zero_model_states[0].param_shapes + avail_numel = fp32_flat_groups[0].numel() * world_size + # Reconstruction protocol: For zero3 we need to zip the partitions together at boundary of each + # param, re-consolidating each param, while dealing with padding if any + + # merge list of dicts, preserving order + param_shapes = {k: v for d in param_shapes for k, v in d.items()} + + if debug: + for i in range(world_size): + print(f"{FP32_FLAT_GROUPS}[{i}].shape={fp32_flat_groups[i].shape}") + + wanted_params = len(param_shapes) + wanted_numel = sum(shape.numel() for shape in param_shapes.values()) + # not asserting if there is a mismatch due to possible padding + avail_numel = fp32_flat_groups[0].numel() * world_size + print(f"Trainable params: Have {avail_numel} numels to process.") + print(f"Trainable params: Need {wanted_numel} numels in {wanted_params} params.") + + # params + # XXX: for huge models that can't fit into the host's RAM we will have to recode this to support + # out-of-core computing solution + offset = 0 + total_numel = 0 + total_params = 0 + for name, shape in param_shapes.items(): + + unpartitioned_numel = shape.numel() + total_numel += unpartitioned_numel + total_params += 1 + + partitioned_numel, partitioned_padding_numel = zero3_partitioned_param_info(unpartitioned_numel, world_size) + + if debug: + print( + f"Trainable params: {total_params} {name} full shape: {shape} partition0 numel={partitioned_numel} partitioned_padding_numel={partitioned_padding_numel}" + ) + + # XXX: memory usage doubles here + state_dict[name] = torch.cat( + tuple(fp32_flat_groups[i].narrow(0, offset, partitioned_numel) for i in range(world_size)), + 0).narrow(0, 0, unpartitioned_numel).view(shape) + offset += partitioned_numel + + offset *= world_size + + # Sanity check + if offset != avail_numel: + raise ValueError(f"consumed {offset} numels out of {avail_numel} - something is wrong") + + print(f"Reconstructed Trainable fp32 state dict with {total_params} params {total_numel} elements") + + +def _get_fp32_state_dict_from_zero3_checkpoint(world_size, fp32_flat_groups, zero_model_states, + exclude_frozen_parameters): + state_dict = OrderedDict() + + # buffers + buffers = zero_model_states[0].buffers + state_dict.update(buffers) + if debug: + print(f"added {len(buffers)} buffers") + + if not exclude_frozen_parameters: + _zero3_merge_frozen_params(state_dict, world_size, zero_model_states) + + _zero3_merge_trainable_params(state_dict, world_size, fp32_flat_groups, zero_model_states) + + # recover shared parameters + for pair in zero_model_states[0].shared_params: + if pair[1] in state_dict: + state_dict[pair[0]] = state_dict[pair[1]] + + return state_dict + + +def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated state_dict that can be loaded with + ``load_state_dict()`` and used for training without DeepSpeed or shared with others, for example + via a model hub. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in 'latest' file. e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + + Returns: + - pytorch ``state_dict`` + + Note: this approach may not work if your application doesn't have sufficient free CPU memory and + you may need to use the offline approach using the ``zero_to_fp32.py`` script that is saved with + the checkpoint. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint + # do the training and checkpoint saving + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu + model = model.cpu() # move to cpu + model.load_state_dict(state_dict) + # submit to model hub or save the model to share with others + + In this example the ``model`` will no longer be usable in the deepspeed context of the same + application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + If you want it all done for you, use ``load_state_dict_from_zero_checkpoint`` instead. + + """ + if tag is None: + latest_path = os.path.join(checkpoint_dir, 'latest') + if os.path.isfile(latest_path): + with open(latest_path, 'r') as fd: + tag = fd.read().strip() + else: + raise ValueError(f"Unable to find 'latest' file at {latest_path}") + + ds_checkpoint_dir = os.path.join(checkpoint_dir, tag) + + if not os.path.isdir(ds_checkpoint_dir): + raise FileNotFoundError(f"Directory '{ds_checkpoint_dir}' doesn't exist") + + return _get_fp32_state_dict_from_zero_checkpoint(ds_checkpoint_dir, exclude_frozen_parameters) + + +def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None, exclude_frozen_parameters=False): + """ + Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be + loaded with ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. + + Args: + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``output_file``: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + - ``exclude_frozen_parameters``: exclude frozen parameters + """ + + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag, exclude_frozen_parameters) + print(f"Saving fp32 state dict to {output_file}") + torch.save(state_dict, output_file) + + +def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None): + """ + 1. Put the provided model to cpu + 2. Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` + 3. Load it into the provided model + + Args: + - ``model``: the model object to update + - ``checkpoint_dir``: path to the desired checkpoint folder. (one that contains the tag-folder, like ``global_step14``) + - ``tag``: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14`` + + Returns: + - ``model`: modified model + + Make sure you have plenty of CPU memory available before you call this function. If you don't + have enough use the ``zero_to_fp32.py`` utility to do the conversion. You will find it + conveniently placed for you in the checkpoint folder. + + A typical usage might be :: + + from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint + model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) + # submit to model hub or save the model to share with others + + Note, that once this was run, the ``model`` will no longer be usable in the deepspeed context + of the same application. i.e. you will need to re-initialize the deepspeed engine, since + ``model.load_state_dict(state_dict)`` will remove all the deepspeed magic from it. + + """ + logger.info(f"Extracting fp32 weights") + state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag) + + logger.info(f"Overwriting model with fp32 weights") + model = model.cpu() + model.load_state_dict(state_dict, strict=False) + + return model + + +if __name__ == "__main__": + + parser = argparse.ArgumentParser() + parser.add_argument("checkpoint_dir", + type=str, + help="path to the desired checkpoint folder, e.g., path/checkpoint-12") + parser.add_argument( + "output_file", + type=str, + help="path to the pytorch fp32 state_dict output file (e.g. path/checkpoint-12/pytorch_model.bin)") + parser.add_argument("-t", + "--tag", + type=str, + default=None, + help="checkpoint tag used as a unique identifier for checkpoint. e.g., global_step1") + parser.add_argument("--exclude_frozen_parameters", action='store_true', help="exclude frozen parameters") + parser.add_argument("-d", "--debug", action='store_true', help="enable debug") + args = parser.parse_args() + + debug = args.debug + + convert_zero_checkpoint_to_fp32_state_dict(args.checkpoint_dir, + args.output_file, + tag=args.tag, + exclude_frozen_parameters=args.exclude_frozen_parameters) diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..8480f3290fe02dcc8ad3f411dce4bb771607f136 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round10_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b2a7a8d47edc80b205a0230dd0d312c16f3cc81aa0b01bd0d21718b6314735bb +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..120369764554a4d6cac58078d8512bf359607829 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round11_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf1dde35905b1ad125edf778aff1f9c3576ae0715c8c9876f665a63efcf18cd +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..64b454d945aa2a842c509c5cebe5ce0706d83a0c --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round12_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb23ed1efc2080025ddd7ca376ff771b3f058afa24e048f5cb3cf99201aa1e03 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..7ad8d07954c41361e73a0c9c7d3b4b46d894b188 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round13_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8caca9dd439d57dbc5054f65d33ebbd7ac694940b2b9dbe5e8cdf68881137a2 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..6f05da1a0577050cea961f470b13842ab5f1d3f9 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round14_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eafca5706ac84f575d7648ee4786543561523e58cea8c01d2214b03d300bb988 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..351a4d2a284717aec90435f50067e1bb7002e9e4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round15_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:341a5a225c4c17a621d4f60caf4c34834000ea74e6f3c8b4f5aaada72ca1f547 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..08cfb19e58b68e297c605b87165d4acd241054e5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round16_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb378347913e2dd7be50fa6f9e3bd9e481f3897c2c9d31c597fe9164efc50bee +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..aeee08397c6123d33648f6703464379a4b7f85fe --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round17_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:70c677ae195453774b8815440fee856ddb773c322681ac422ebc6e8abd5a58d7 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..f9935621fdc6da9e4abad379d97303b5ebf4f8f5 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round18_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:393c3d5a15b40b2ff049c254f5b09f25d226b0971827b4e7460ca58f9ff7b76f +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..2766cf1a4da1390df01ea4146b4a3a7a06e1434e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round19_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5dc77db94ac5b4746eebf2c850a9425da9a425a71108790d6b7f5f2654b614a8 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..34bfc5121e29fc587c3f574c17d8006189d3021a --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round1_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47a94b88b812f21101345aae20a53cc10e0b346660c2cef669c77b201d06d508 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..4be93ba48f85d89c8cbd0e7426b79a10f51bc786 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round20_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20b9a2e67fdcb4051cebf52cba2ed0d47facb4dc69170f35a3483e3601bf25a4 +size 704650017 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..b42d717305b30cd186c1c408b8719020c0658a5d --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round2_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9fd7458877767da0a36b80344cd6ec7eae8fbc8e11d07f8bfce3b4d5f5f239a9 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..01842d1316f21c3326c5735c6257b499da7ca458 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round3_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e537f056156082569ae0819abd06beb04b46581d0c63843bd68926e0ac2bff0e +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..3af9fe115499486166c3ef6d0a7bd27661ac08b4 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round4_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:92a9649d6997f84508e5f566393e91e42e03a31aca55878bacd7ffd2e134bbfc +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..abd1e557c23a8769870b0beedaf7df93836a101e --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round5_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c7520bace4f501771d879621facd1b44f8db1ad7385a653e1535c7642465e67d +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..ebc846b7a5c587ced387cca85c4065e921c8cbe3 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round6_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1951a910d1b7587459db2110e8ac43cbc906a34267372eb21d2244a900eaec46 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..5eb1370d3fa3177cb4fb710882f4b7eb706817e8 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round7_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:03f40a0df2437812b9e6de4e5d4e30014acb71344efa8132f19e1fc0f5925c82 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..385c599585a0973f240a027d28e985ca72d6b721 --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round8_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9f3ef6d8247bc37b017895c9e5bbaf32f2db500f94f66b4ebade5a4a5ba7a271 +size 704649992 diff --git a/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth new file mode 100644 index 0000000000000000000000000000000000000000..219adb8927937a4597cd82efbf7ebb4c8a7519ad --- /dev/null +++ b/client_states_feddat_feddualMultipqfullfreeze_homoAgg_moe_iter100_round5_hetero/round9_task_vector_local_weights.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:277fa5f0607d993e80153171e4760f999c4ee6a65452ea62694d0b9b03818f22 +size 704649992