Como manipular gradientes de cliente em tensorflow federated sgd
Estou seguindo este tutorial para começar a usar tensorflow federado. Meu objetivo é executar o sgd federado (não avg federado) com algumas manipulações nos valores de gradiente do cliente antes de serem enviados ao servidor.
Antes de prosseguir, para reiterar brevemente o processo sgd federado, para cada vez os clientes enviarão seus gradientes computados (pesos não atualizados) para o servidor, o servidor os agrega e transmite o modelo atualizado para os clientes.
Agora, pelo que reuni até agora, posso usar a função em build_federated_sgd_processvez da build_federated_averaging_processdo tutorial mencionado para executar sgd federado da maneira descrita acima.
Estou perdido, preciso cortar os gradientes do cliente e adicionar algum ruído a eles (gerado de forma independente para cada valor de gradiente) antes de enviar os gradientes para o servidor e não tenho certeza de como fazer isso. Gerar o ruído é bastante simples, mas qual função devo modificar / implementar para poder aplicar o ruído aos gradientes?
Respostas
build_federated_sgd_processestá totalmente enlatado; ele é realmente projetado para servir como uma implementação de referência, não como um ponto de extensibilidade.
Eu acredito que o que você está procurando é a função que build_federated_sgd_processchama sob o hoos tff.learning.framework.build_model_delta_optimizer_process,. Esta função permite que você forneça seu próprio mapeamento de uma função de modelo (ou seja, um zero-arg chamável que retorna a tff.learning.Model) para a tff.learning.framework.ClientDeltaFn.
Sua ClientDeltaFnaparência seria algo como:
@tf.function
def _clip_and_noise(grads):
return ...
class ClippedGradClientDeltaFn(tff.learning.framework.ClientDeltaFn)
def __init__(self, model, ...):
self._model = model
...
@tf.function
def __call__(dataset, weights):
# Compute gradients grads
return _clip_and_noise(grads)
E você seria capaz de construir um tff.templates.IterativeProcesschamando:
def clipped_sgd(model_fn: Callable[[], model_lib.Model]) -> ClippedGradClientDeltaFn:
return ClippedGradClientDeltaFn(
model_fn(),
...)
iterproc = optimizer_utils.build_model_delta_optimizer_process(
model_fn, model_to_client_delta_fn=clipped_sgd, ...)
como mais ou menos no corpo de build_federated_sgd_process.
Parece-me que você está interessado em privacidade diferencial; O TFF é, na verdade, projetado para compor com privacidade diferencial geralmente por meio dos processos de agregação, em vez de escrever diferentes atualizações do cliente, embora esta seja certamente uma abordagem. Consulte as dicas da TFF para documentação de pesquisa sobre maneiras idiomáticas de conectar privacidade diferencial à TFF.