Qiskit混合神经网络中@staticmethod在backward函数的作用解析
Qiskit教程中反向传播函数@staticmethod的作用解析
先明确@staticmethod的基础作用
@staticmethod是Python的装饰器,用来将类中的方法标记为静态方法:
- 它隶属于类本身,无需创建类的实例就能调用;
- 不会自动接收
self(实例引用)作为第一个参数,只能使用传入的参数或类的静态属性。
为什么PyTorch自定义反向传播要加这个装饰器
你贴的这段代码是在自定义PyTorch的Function类(用于实现量子层的自定义前向/反向传播逻辑),而PyTorch的官方规则明确要求:自定义Function的backward方法必须是静态方法。
- 经典神经网络用
torch.nn.Module搭建时,反向传播是框架自动生成的,不需要手动编写backward方法,所以你看不到这个装饰器; - 但量子层属于自定义算子,必须继承
torch.autograd.Function并手动实现forward和backward,此时backward必须用@staticmethod标记——因为PyTorch在执行反向传播时,是直接通过类来调用该方法,而非类的实例。
结合你提供的代码看具体适配性
这段backward的所有操作逻辑:
- 从
ctx(PyTorch的上下文对象,用于存储前向传播的中间结果)中取出saved_tensors和shift参数; - 计算参数偏移后的量子电路期望;
- 最终计算梯度并返回。
所有需要的变量都通过传入的参数获取,完全不需要访问类的实例属性,用静态方法刚好匹配这种场景,也完全符合PyTorch的要求。
附上你提供的代码:
@staticmethod def backward(ctx, grad_output): """ Backward pass computation """ input, expectation_z = ctx.saved_tensors input_list = np.array(input.tolist()) shift_right = input_list + np.ones(input_list.shape) * ctx.shift shift_left = input_list - np.ones(input_list.shape) * ctx.shift gradients = [] for i in range(len(input_list)): expectation_right = ctx.quantum_circuit.run(shift_right[i]) expectation_left = ctx.quantum_circuit.run(shift_left[i]) gradient = torch.tensor([expectation_right]) - torch.tensor([expectation_left]) gradients.append(gradient) gradients = np.array([gradients]).T return torch.tensor([gradients]).float() * grad_output.float(), None, None
内容的提问来源于stack exchange,提问作者mchaudh4
相关产品推荐
相关产品推荐

