РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
323 of 824 menu

Атрибут trainable_variables

Атрибут trainable_variables класса Module возвращает список обучаемых переменных, принадлежащих модулю. К обучаемым переменным относятся веса и смещения слоев, которые обновляются в процессе обучения. Атрибут не принимает параметров и доступен для чтения у любого объекта, унаследованного от Module, в том числе у моделей и слоев Keras.

Синтаксис

module.trainable_variables

Пример

Давайте создадим простой слой Dense и посмотрим на его обучаемые переменные:

import tensorflow as tf tf.random.set_seed(0) layer = tf.keras.layers.Dense(3, input_shape=(2,)) res = layer.trainable_variables print(res)

Результат выполнения кода:

[<tf.Variable 'kernel:0' shape=(2, 3) dtype=float32, numpy= array([[...]], dtype=float32)>, <tf.Variable 'bias:0' shape=(3,) dtype=float32, numpy=array([...], dtype=float32)>]

Пример

Давайте создадим модель из двух слоев и выведем количество и формы обучаемых переменных:

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(4, input_shape=(3,)), tf.keras.layers.Dense(2) ]) vars = model.trainable_variables print("count:", len(vars)) for v in vars: print(v.name, v.shape)

Результат выполнения кода:

count: 4 dense/kernel:0 (3, 4) dense/bias:0 (4,) dense_1/kernel:0 (4, 2) dense_1/bias:0 (2,)

Пример

Давайте обнулим обучаемые переменные модели и проверим их значения:

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(2, input_shape=(3,)) ]) for v in model.trainable_variables: v.assign(tf.zeros_like(v)) res = model.trainable_variables print(res)

Результат выполнения кода:

[<tf.Variable 'kernel:0' shape=(3, 2) dtype=float32, numpy= array([[0., 0.], [0., 0.], [0., 0.]], dtype=float32)>, <tf.Variable 'bias:0' shape=(2,) dtype=float32, numpy=array([0., 0.], dtype=float32)>]

Смотрите также

  • класс Module,
    который является базовым классом для моделей и слоев
  • атрибут variables,
    который возвращает все переменные модуля
  • метод call,
    который определяет прямой проход модуля
  • метод __call__,
    который вызывает модуль как функцию
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить