【问题标题】:Why does my model learn with Ragged Tensors but not Dense Tensors?为什么我的模型使用不规则张量而不是密集张量学习?
【发布时间】:2021-05-03 14:53:50
【问题描述】:

我有一串遵循“语法”的字母。我的训练集上也有关于字符串是否遵循“语法”的布尔标签。基本上,我的模型试图学习确定一串字母是否符合规则。这是一个相当简单的问题(我是从教科书中得到的)。

我正在像这样生成我的数据集:

def generate_dataset(size):
    good_strings = [string_to_ids(generate_string(embedded_reber_grammar))
                    for _ in range(size // 2)]
    bad_strings = [string_to_ids(generate_corrupted_string(embedded_reber_grammar))
                   for _ in range(size - size // 2)]
    all_strings = good_strings + bad_strings
    X = tf.ragged.constant(all_strings, ragged_rank=1)

    # X = X.to_tensor(default_value=0)

    y = np.array([[1.] for _ in range(len(good_strings))] +
                 [[0.] for _ in range(len(bad_strings))])
    return X, y

注意X = X.to_tensor(default_value=0) 这一行。如果这条线被注释掉,我的模型就学得很好。但是,如果它没有被注释掉,它就无法学习,验证集的表现与机会(50-50)相同。

这是我的实际模型:

np.random.seed(42)
tf.random.set_seed(42)

embedding_size = 5

model = keras.models.Sequential([
    keras.layers.InputLayer(input_shape=[None], dtype=tf.int32, ragged=True),
    keras.layers.Embedding(input_dim=len(POSSIBLE_CHARS) + 1, output_dim=embedding_size),
    keras.layers.GRU(30),
    keras.layers.Dense(1, activation="sigmoid")
])
optimizer = keras.optimizers.SGD(lr=0.02, momentum = 0.95, nesterov=True)
model.compile(loss="binary_crossentropy", optimizer=optimizer, metrics=["accuracy"])
history = model.fit(X_train, y_train, epochs=5, validation_data=(X_valid, y_valid))

我使用0 作为密集张量的默认值。 strings_to_ids 没有对任何值使用 0,而是从 1 开始。此外,当我切换到使用密集张量时,我将 ragged=True 更改为 False. 我不知道为什么使用密集张量会导致模型失败,因为我之前在类似的练习中使用过密集张量。

有关更多详细信息,请参阅书中的解决方案 (exercise 8) 或我自己的 colab notebook

【问题讨论】:

    标签: tensorflow machine-learning keras machine-learning-model


    【解决方案1】:

    所以答案是密集张量的形状在训练集和验证集上是不同的。这是因为两个集合之间的最长序列长度不同(与测试集相同)。

    【讨论】:

      猜你喜欢
      • 2019-11-20
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2023-03-12
      • 2020-11-18
      • 1970-01-01
      • 2018-04-29
      相关资源
      最近更新 更多