pip install tensorflow python import tensorflow as tf dataset = ... def preprocess_fn(data): ... dataset = dataset.map(preprocess_fn) train_dataset = dataset.take(train_size) test_dataset = dataset.skip(train_size) train_dataset = train_dataset.shuffle(buffer_size).batch(batch_size) test_dataset = test_dataset.batch(batch_size) train_iterator = train_dataset.make_initializable_iterator() test_iterator = test_dataset.make_initializable_iterator() train_data = train_iterator.get_next() test_data = test_iterator.get_next() model_input = ... ... ... ... ... python import tensorflow as tf (x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data() x_train = x_train.astype('float32') / 255.0 x_test = x_test.astype('float32') / 255.0 model = tf.keras.Sequential([...]) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) model.fit(x_train, y_train, epochs=5, batch_size=64) python import tensorflow as tf model = tf.keras.Sequential([...]) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) model.fit(train_data, epochs=5, batch_size=64) loss, accuracy = model.evaluate(test_data) python import tensorflow as tf cluster_spec = tf.train.ClusterSpec({ 'worker': ['worker0:port', 'worker1:port', ...], 'ps': ['ps0:port', 'ps1:port', ...] }) server = tf.train.Server(cluster_spec, job_name='worker', task_index=0) with tf.device('/job:worker/task:0'): ...


上一篇:
下一篇:
切换中文