Logo Lanfrica

VityaVitalich/MeritOpt

Domain:

natural language processing

Record type:

software
Creator:
Vit
Host:
[EMNLP 2024] Low-Resource Machine Translation through the Lens of Personalized Federated Learning # [EMNLP 2024] Low-Resource Machine Translation through the Lens of Personalized Federated Learning This repository contains code for paper Low-Resource Machine Translation through the Lens of Personalized Federated Learning ## Using optimizer The optimizer could be found in ```pipeline_src/optimizers.py```. To add this into your code you just need to import the optimizer and correctly provide the losses to it during training. Below is the example of training with Indonesian and Javanese languages. Our code also requires accelerate to run. ```python from pipeline_src.optimizers import MeritFedParallelMD from accelerate import Accelerator # Init the accelerator accelerator = Accelerator() config = { 'lr': , 'npeers': , 'mdlr_': , 'mdniters_' = , 'drop_threshold' = } device = model = # wrap the model with accelerator model = accelerator.prepare_model(model) weight_name_map = # for example {0: indonesian, 1: javanese} train_loader, val_loader = # wrap dataloaders with accelerator train_loader = accelerator.prepare(train_loader) val_loader = accelerator.prepare(val_loader) optimizer = MeritFedA( model.parameters(), config, val_loader=val_loader, model=model, accelerator=accelerator ) # During training we need to register each worker grad # First we calculate loss on the Indonesian Data w_id = 0 # We have set id 0 to indonesian output = model.forward(indonesian_input) loss = output["loss"] loss.backward() # We step with providing id of data, model and validation loader to perform auxiliary optimization # double optimizer class since first is wrapper of accelerate optimizer.optimizer.register_worker_grad(w_id) # we perform the zero grad here optimizer.zero_grad() # Next we do the same with second language, javanese in our example w_id = 1 # Javanese has id = 1 output = model.forward(javanese_input) loss = output["loss"] loss.backward() # We step the same way but with new id optimizer.optimizer.register_worker_grad(w_id) # WE DO NOT PERFORM ZERO GRAD AT LA …