[Embedding] Check the sharded property of tf.train.Saver. (#996)

Signed-off-by: chenbangduo.cbd <chenbangduo.cbd@alibaba-inc.com>
This commit is contained in:
Chen Bangduo 2024-05-23 12:00:02 +08:00 committed by GitHub
parent 93c69ad957
commit 9e30ab604a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
22 changed files with 76 additions and 71 deletions

View File

@ -612,10 +612,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -527,10 +527,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -594,10 +594,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -610,10 +610,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -472,10 +472,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -776,10 +776,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -594,10 +594,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -507,10 +507,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -478,10 +478,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -534,10 +534,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=train_steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -529,10 +529,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -522,10 +522,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -523,10 +523,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -592,10 +592,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -427,10 +427,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=train_steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -543,10 +543,9 @@ def train(sess_config,
hooks = []
hooks.extend(input_hooks)
sharded_saver = tf_config != None
scaffold = tf.train.Scaffold(
local_init_op=tf.group(tf.local_variables_initializer(), data_init_op),
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=sharded_saver))
saver=tf.train.Saver(max_to_keep=args.keep_checkpoint_max, sharded=True))
stop_hook = tf.train.StopAtStepHook(last_step=steps)
log_hook = tf.train.LoggingTensorHook(

View File

@ -7527,7 +7527,7 @@ class EmbeddingColumnTest(test.TestCase):
opt = ftrl.FtrlOptimizer(0.1, l1_regularization_strength=2.0, l2_regularization_strength=0.00001)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables_lib.global_variables_initializer()
with self.test_session() as sess:
sess.run(ops.get_collection(ops.GraphKeys.EV_INIT_VAR_OPS))
@ -7758,7 +7758,7 @@ class EmbeddingColumnTest(test.TestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
init = variables_lib.global_variables_initializer()
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
@test_util.run_deprecated_v1
def testEmbeddingVariableForInt32ID(self):
@ -7783,7 +7783,7 @@ class EmbeddingColumnTest(test.TestCase):
opt = ftrl.FtrlOptimizer(0.1, l1_regularization_strength=2.0, l2_regularization_strength=0.00001)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables_lib.global_variables_initializer()
with self.test_session() as sess:
sess.run(ops.get_collection(ops.GraphKeys.EV_INIT_VAR_OPS))

View File

@ -63,7 +63,8 @@ class EmbeddingVariableGpuTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
graph = ops.get_default_graph()
meta_graph_def = saver_module.export_meta_graph()
saver = saver_module.Saver(sharded=True)
meta_graph_def = saver_module.export_meta_graph(saver_def=saver.as_saver_def())
ops.reset_default_graph()
with self.test_session() as sess:
res = saver_module.import_meta_graph(meta_graph_def)
@ -748,7 +749,7 @@ class EmbeddingVariableGpuTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, global_step=gs)
init = variables.global_variables_initializer()
saver = saver = saver_module.Saver()
saver = saver = saver_module.Saver(sharded=True)
checkpoint_directory = self.get_temp_dir()
model_path = os.path.join(checkpoint_directory, "model.ckpt")
with self.test_session() as sess:
@ -816,7 +817,7 @@ class EmbeddingVariableGpuTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
graph = ops.get_default_graph()
with self.test_session(graph = graph) as sess:
saver.restore(sess, os.path.join(checkpoint_directory, "model.ckpt-12345"))

View File

@ -162,7 +162,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = self._CreateOptimizer(optimizer)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -194,7 +194,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = self._CreateOptimizer(optimizer)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -232,7 +232,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, global_step=gs)
init = variables.global_variables_initializer()
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
model_path = os.path.join(checkpoint_directory, "model.ckpt")
with self.test_session() as sess:
sess.run([init])
@ -269,7 +269,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -313,7 +313,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = gradient_descent.GradientDescentOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -387,7 +387,8 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
graph = ops.get_default_graph()
meta_graph_def = saver_module.export_meta_graph()
saver = saver_module.Saver(sharded=True)
meta_graph_def = saver_module.export_meta_graph(saver_def=saver.as_saver_def())
ops.reset_default_graph()
with self.test_session() as sess:
res = saver_module.import_meta_graph(meta_graph_def)
@ -406,7 +407,8 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
graph = ops.get_default_graph()
meta_graph_def = saver_module.export_meta_graph()
saver = saver_module.Saver(sharded=True)
meta_graph_def = saver_module.export_meta_graph(saver_def=saver.as_saver_def())
ops.reset_default_graph()
with self.test_session() as sess:
res = saver_module.import_meta_graph(meta_graph_def)
@ -450,7 +452,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adam.AdamOptimizer(0.01)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session() as sess:
sess.run(ops.get_collection(ops.GraphKeys.EV_INIT_VAR_OPS))
@ -643,7 +645,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = ftrl.FtrlOptimizer(0.1, l1_regularization_strength=2.0, l2_regularization_strength=0.00001)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session() as sess:
sess.run(ops.get_collection(ops.GraphKeys.EV_INIT_VAR_OPS))
@ -682,7 +684,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, global_step=gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session() as sess:
sess.run([init])
@ -720,7 +722,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = ftrl.FtrlOptimizer(0.1, l1_regularization_strength=2.0, l2_regularization_strength=0.00001)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session() as sess:
sess.run(ops.get_collection(ops.GraphKeys.EV_INIT_VAR_OPS))
@ -1534,7 +1536,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
init = variables.global_variables_initializer()
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
model_path = os.path.join(checkpoint_directory, "model.ckpt")
with self.test_session() as sess:
sess.run([init])
@ -1567,7 +1569,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = ftrl.FtrlOptimizer(0.1, l1_regularization_strength=2.0, l2_regularization_strength=0.00001)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session() as sess:
sess.run(ops.get_collection(ops.GraphKeys.EV_INIT_VAR_OPS))
@ -1724,7 +1726,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, global_step=gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -1778,7 +1780,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, global_step=gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -1849,7 +1851,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, global_step=gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -1923,7 +1925,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -1963,7 +1965,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -2278,7 +2280,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = ftrl.FtrlOptimizer(0.1, l1_regularization_strength=2.0, l2_regularization_strength=0.00001)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
def testSaveV3(self):
print("testSaveV3")
@ -2295,7 +2297,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, global_step=gs)
init = variables.global_variables_initializer()
saver = saver = saver_module.Saver()
saver = saver = saver_module.Saver(sharded=True)
checkpoint_directory = self.get_temp_dir()
model_path = os.path.join(checkpoint_directory, "model.ckpt")
with self.test_session() as sess:
@ -2326,7 +2328,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -2359,7 +2361,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
model_path = os.path.join(checkpoint_directory,
"model1.ckpt")
@ -2390,7 +2392,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad.AdagradOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v, gs)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session() as sess:
sess.run([init])
@ -2412,7 +2414,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
emb = embedding_ops.embedding_lookup(emb_var, ids)
tires = kv_variable_ops.lookup_tier(emb_var,
math_ops.cast([1,2,3,4], dtypes.int64))
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
graph = ops.get_default_graph()
with self.test_session(graph = graph) as sess:
saver.restore(sess, os.path.join(checkpoint_directory, "model.ckpt"))
@ -2784,7 +2786,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
init = variables.global_variables_initializer()
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
with self.test_session() as sess:
result = sess.run(var._is_initialized_op)
self.assertEqual(False, result)
@ -2806,7 +2808,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad_decay.AdagradDecayOptimizer(0.1, gs)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session(graph=g) as sess:
sess.run([init])
@ -2823,7 +2825,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad_decay.AdagradDecayOptimizer(0.1, gs)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session(graph=g) as sess:
result = sess.run(var._is_initialized_op)
@ -2860,7 +2862,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad_decay.AdagradDecayOptimizer(0.1, gs)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session(graph=g) as sess:
sess.run([init])
@ -2893,7 +2895,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = adagrad_decay.AdagradDecayOptimizer(0.1, gs)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session(graph=g) as sess:
sess.run([init])
@ -2929,7 +2931,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = gradient_descent.GradientDescentOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session(graph=g) as sess:
sess.run([init])
@ -2964,7 +2966,7 @@ class EmbeddingVariableTest(test_util.TensorFlowTestCase):
opt = gradient_descent.GradientDescentOptimizer(0.1)
g_v = opt.compute_gradients(loss)
train_op = opt.apply_gradients(g_v)
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
with self.test_session(graph=g) as sess:
sess.run([init])

View File

@ -75,7 +75,7 @@ class IncrSaveRestoreTest(test_util.TensorFlowTestCase):
emb = embedding_ops.embedding_lookup(var, math_ops.cast([0,1,2,5,6,7], dtypes.int64))
with ops.device("/device:CPU:0"):
apply_incr = gen_io_ops.record_sparse_indices(math_ops.cast([0,1,2,5,6,7], dtypes.int64), "var_ev1")
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
ev_var_name = "var_ev1"
incr_save_op = gen_io_ops.incr_save(incr_ckpt_path, [ev_var_name], [], [True],[var.handle])
@ -178,7 +178,7 @@ class IncrSaveRestoreTest(test_util.TensorFlowTestCase):
activate_op = gen_io_ops. activate_sparse_recorder(["var_ev1","var_norm1"])
saver = saver_module.Saver()
saver = saver_module.Saver(sharded=True)
init = variables.global_variables_initializer()
incr_save_op = gen_io_ops.incr_save(incr_ckpt_path, ["var_norm1", "var_ev1"], [], [True, True], [var_norm, var_ev.handle])
@ -445,6 +445,7 @@ class IncrSaveRestoreTest(test_util.TensorFlowTestCase):
variable_scope.get_variable('var', shape=[100], use_resource=False)
variable_scope.get_embedding_variable('ev', embedding_dim=100)
saver = saver_module.Saver(
sharded=True,
save_relative_paths=True,
incremental_save_restore=True,
)

View File

@ -1071,10 +1071,14 @@ class Saver(object):
# pylint: disable=protected-access
self._var_list = variables._all_saveable_objects()
from tensorflow.python.ops import hash_table
from tensorflow.python.ops import kv_variable_ops
if isinstance(self._var_list, dict):
ev = {}
ht = {}
lst = {}
for name, x in self._var_list.items():
if isinstance(x, kv_variable_ops.EmbeddingVariable):
ev[name] = x
if isinstance(x, hash_table.HashTable):
if x.hash_table not in ht:
ht[x.hash_table] = [x]
@ -1084,15 +1088,20 @@ class Saver(object):
lst[name] = BloomFilterSaveable(x)
else:
lst[name] = x
if len(ev) != 0 and not self._sharded:
raise ValueError("EmbeddingVariable can only use sharded saver")
if len(ht) != 0 and not self._sharded:
raise ValueError("HashTable can only use sharded saver")
for x, y in ht.items():
lst[x.name] = HashTableSaveable(y)
self._var_list = lst
else:
ev = []
ht = {}
lst = []
for x in self._var_list:
if isinstance(x, kv_variable_ops.EmbeddingVariable):
ev.append(x)
if isinstance(x, hash_table.HashTable):
if x.hash_table not in ht:
ht[x.hash_table] = [x]
@ -1102,6 +1111,8 @@ class Saver(object):
lst.append(BloomFilterSaveable(x))
else:
lst.append(x)
if len(ev) != 0 and not self._sharded:
raise ValueError("EmbeddingVariable can only use sharded saver")
if len(ht) != 0 and not self._sharded:
raise ValueError("HashTable can only use sharded saver")
for x, y in ht.items():

View File

@ -852,6 +852,12 @@ class SaverTest(test.TestCase):
for orig, restored in zip(orig_vals, restored_vals):
self.assertAllEqual(orig, restored)
def testEnableSaverShardedWhenUseEmbeddingVariable(self):
with ops_lib.Graph().as_default():
emb_var = \
variable_scope.get_embedding_variable(name="emb_var", embedding_dim=64)
with self.assertRaisesRegexp(ValueError, "EmbeddingVariable"):
saver_module.Saver([emb_var], sharded=False)
class SaveRestoreShardedTest(test.TestCase):