You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Python类中调用self()的作用?Keras Transformer代码解析

Keras Transformer中self()调用的解析

我在尝试用Keras实现Transformer时遇到如下代码:

class Transformer(keras.Model):
    def __init__(
        self,
        num_hid=64,
        num_head=2,
        num_feed_forward=128,
        source_maxlen=100,
        target_maxlen=100,
        num_layers_enc=4,
        num_layers_dec=1,
        num_classes=60,
    ):
        super().__init__()
    
    ...

    def train_step(self, batch):

        ....

        with tf.GradientTape() as tape:
            preds = self([source, dec_input]) #疑问行
            one_hot = tf.one_hot(dec_target, depth=self.num_classes)
            mask = tf.math.logical_not(tf.math.equal(dec_target, pad_token_idx))
            loss = self.compiled_loss(one_hot, preds, sample_weight=mask)

核心疑问

  1. Transformer类的train_step方法中调用self()是何含义?它在初始化哪个类?是GradientTape还是Model类?
  2. Python允许在子类方法中用self初始化父类吗?
  3. 为什么测试代码报错,无法在对象内部用self()创建新对象?

问题解析

1. 原代码中self([source, dec_input])的含义

这段代码不是初始化任何类,也和GradientTape无关——GradientTape仅用于记录梯度计算的上下文,方便后续反向传播求导。

Transformer类继承自keras.Model,而Keras的Model类已经实现了__call__魔法方法。当调用self([source, dec_input])时,本质是触发Model的__call__方法,最终会执行你在Transformer类中定义的call方法(若有自定义实现),完成模型的前向传播计算:输入source和dec_input,输出预测结果preds。

简单来说,这和你平时调用model(input_data)生成预测的逻辑完全一致,就是用当前模型实例做一次前向计算。

2. Python中用self()初始化父类?

这是对代码逻辑的误解——self()根本不是初始化父类的操作。正常情况下,普通Python类如果未实现__call__方法,直接调用self()会报错,就像你的测试代码那样。只有实现了__call__魔法方法的对象,才能被当作函数调用。

3. 测试代码报错的原因

你的测试代码:

class A:
    def __init__(self, x=None, y=None):
        print("A:",x,y)

class B(A):
    def __init__(self):
        super().__init__()
    def call(self, x, y):
        c = self(x,y)

b = B()
b.call(1,2)

报错TypeError: 'B' object is not callable的核心原因:

  • 你定义的call是普通方法,并非Python的魔法方法__call__,Python不会将其识别为对象可调用的逻辑。
  • 类A和B都未实现__call__魔法方法,因此self(x,y)这种调用方式不被允许——Python无法将B的实例当作函数执行。

如果要让测试代码正常运行,需把普通方法call改为魔法方法__call__:

class A:
    def __init__(self, x=None, y=None):
        print("A:",x,y)

class B(A):
    def __init__(self):
        super().__init__()
    def __call__(self, x, y):
        print("B call:", x, y)
        return x + y

b = B()
b(1,2) # 直接调用实例触发__call__方法

内容的提问来源于stack exchange,提问作者Akanksha Makkar

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 17:32:41