Como o código de multiprocessamento seleciona a função do mapa?
Estou escrevendo um utilitário de pesquisa de grade e estou tentando usar o multiprocessamento para acelerar o cálculo. Tenho uma função objetivo que interage com uma grande classe que não posso conservar devido a restrições de memória (só posso conservar atributos relevantes da classe).
import pickle
from multiprocessing import Pool
class TestClass:
def __init__(self):
self.param = 10
def __getstate__(self):
raise RuntimeError("don't you dare pickle me!")
def __setstate__(self, state):
raise RuntimeError("don't you dare pickle me!")
def loss(self, ext_param):
return self.param*ext_param
if __name__ == '__main__':
test_instance = TestClass()
def objective_function(param):
return test_instance.loss(param)
with Pool(4) as p:
result = p.map(objective_function, range(20))
print(result)
No exemplo de brinquedo a seguir, eu esperava durante a decapagem da função objetivo, que test_instance também teria que ser decapado, lançando assim RuntimeError (devido à exceção lançada em __getstate__). No entanto, isso não acontece e o código é executado sem problemas.
Portanto, minha pergunta é - o que está sendo conservado aqui exatamente? E se test_instance não é conservado, então como ele é reconstruído em processos individuais?
Respostas
No windows + python3.8, não consegui executar o código original que definiu test_instance e goal_function como variável local para principal, erro como abaixo
AttributeError: Can't get attribute 'objective_function' on <module '__mp_main' from 'xxx.py'>
Mudei a definição de goal_function e a inicialização de test_instance para o escopo global, funciona bem como você mencionou. No entanto, a partir disso , parece que as variáveis globais foram inicializadas novamente para processos diferentes ao invés de conservadas / não conservadas.
Por fim, alterei seu código conforme a seguir e ele acionou o erro que você esperava.
test_instance1 = TestClass()
test_instance2 = TestClass()
with Pool(4) as p:
result = p.map(objective_function, [test_instance1, test_instance2])
print(result)
Portanto, os prameters em p.map realmente fazem pickle / unpickled.
Ok, com a ajuda de Wilson e mais algumas pesquisas, consegui responder minha própria pergunta. Vou inserir o código modificado acima para ajudar na explicação:
import pickle
from multiprocessing import Pool, current_process
class TestClass:
def __init__(self):
self.param = 0
def __getstate__(self):
raise RuntimeError("don't you dare pickle me!")
def __setstate__(self, state):
raise RuntimeError("don't you dare pickle me!")
def loss(self, ext_param):
self.param += 1
print(f"{current_process().pid}: {hex(id(self))}: {self.param}: {ext_param} ")
return f"{self.param}_{ext_param}"
def objective_function(param):
return test_instance.loss(param)
if __name__ == '__main__':
test_instance = TestClass()
print(hex(id(test_instance)))
print('objective_function' in globals()) # this returns True on my MacOS+python3.7
with Pool(2) as p:
result = p.map(objective_function, range(6))
print(result)
print(test_instance.param)
# ---- RUN RESULTS BELOW ----
# 0x7f987b955e48
# True
# 10484: 0x7f987b955e48: 1: 0
# 10485: 0x7f987b955e48: 1: 1
# 10484: 0x7f987b955e48: 2: 2
# 10485: 0x7f987b955e48: 2: 3
# 10484: 0x7f987b955e48: 3: 4
# 10485: 0x7f987b955e48: 3: 5
# ['1_0', '1_1', '2_2', '2_3', '3_4', '3_5']
# 0
Como Wilson sugeriu corretamente, a única coisa que fica bloqueada durante o p.map são os próprios parâmetros e não a função objetivo, no entanto, isso não é reinicializado, mas copiado, junto com o test_instance durante o os.fork()processo que acontece em algum lugar na inicialização do Pool. Você pode ver que, embora dentro de cada processo os test_instance.paramvalores sejam independentes uns dos outros, eles compartilham a mesma memória virtual que a instância original da classe antes da bifurcação (um exemplo de processos diferentes compartilhando a mesma memória virtual pode ser visto aqui ).
Quanto à solução da questão inicial, acredito que a única maneira de resolver este problema adequadamente é distribuir os parâmetros necessários na memória compartilhada ou gerenciador de memória.