@@ -85,85 +85,87 @@ def dataset_fn(input_context):
8585
8686
8787class KerasPremadeModelsTest (tf .test .TestCase , parameterized .TestCase ):
88- @tf .__internal__ .distribute .combinations .generate (
88+ @tf .__internal__ .distribute .combinations .generate (
8989 strategy_combinations_eager_data_fn ()
9090 )
91- def test_linear_model (self , distribution , use_dataset_creator , data_fn ):
92- if (not use_dataset_creator ) and isinstance (
91+ def test_linear_model (self , distribution , use_dataset_creator , data_fn ):
92+ if (not use_dataset_creator ) and isinstance (
9393 distribution , tf .distribute .experimental .ParameterServerStrategy
9494 ):
95- self .skipTest (
95+ self .skipTest (
9696 "Parameter Server strategy requires dataset creator to be used "
9797 "in model.fit."
9898 )
99- if (
99+ if (
100100 not tf .__internal__ .tf2 .enabled ()
101101 and use_dataset_creator
102102 and isinstance (
103103 distribution , tf .distribute .experimental .ParameterServerStrategy
104104 )
105105 ):
106- self .skipTest (
106+ self .skipTest (
107107 "Parameter Server strategy with dataset creator needs to be "
108108 "run when eager execution is enabled."
109109 )
110- with distribution .scope ():
111- model = linear .LinearModel ()
112- opt = gradient_descent .SGD (learning_rate = 0.1 )
113- model .compile (opt , "mse" )
114- if use_dataset_creator :
115- x = dataset_creator .DatasetCreator (dataset_fn )
116- hist = model .fit (x , epochs = 3 , steps_per_epoch = INPUT_SIZE )
117- else :
118- if data_fn == "numpy" :
119- inputs , output = get_numpy ()
120- hist = model .fit (inputs , output , epochs = 3 )
121- else :
122- hist = model .fit (get_dataset (), epochs = 3 )
123- self .assertLess (hist .history ["loss" ][2 ], 0.2 )
124-
125- @tf .__internal__ .distribute .combinations .generate (
110+ with distribution .scope ():
111+ model = linear .LinearModel ()
112+ opt = gradient_descent .SGD (learning_rate = 0.1 )
113+ model .compile (opt , "mse" )
114+ if use_dataset_creator :
115+ x = dataset_creator .DatasetCreator (dataset_fn )
116+ hist = model .fit (x , epochs = 3 , steps_per_epoch = INPUT_SIZE )
117+ else :
118+ if data_fn == "numpy" :
119+ inputs , output = get_numpy ()
120+ hist = model .fit (inputs , output , epochs = 3 )
121+ else :
122+ hist = model .fit (get_dataset (), epochs = 3 , steps_per_epoch = INPUT_SIZE )
123+ self .assertLess (hist .history ["loss" ][2 ], 0.2 )
124+
125+ @tf .__internal__ .distribute .combinations .generate (
126126 strategy_combinations_eager_data_fn ()
127127 )
128- def test_wide_deep_model (self , distribution , use_dataset_creator , data_fn ):
129- if (not use_dataset_creator ) and isinstance (
128+ def test_wide_deep_model (self , distribution , use_dataset_creator , data_fn ):
129+ if (not use_dataset_creator ) and isinstance (
130130 distribution , tf .distribute .experimental .ParameterServerStrategy
131131 ):
132- self .skipTest (
132+ self .skipTest (
133133 "Parameter Server strategy requires dataset creator to be used "
134134 "in model.fit."
135135 )
136- if (
136+ if (
137137 not tf .__internal__ .tf2 .enabled ()
138138 and use_dataset_creator
139139 and isinstance (
140140 distribution , tf .distribute .experimental .ParameterServerStrategy
141141 )
142142 ):
143- self .skipTest (
143+ self .skipTest (
144144 "Parameter Server strategy with dataset creator needs to be "
145145 "run when eager execution is enabled."
146146 )
147- with distribution .scope ():
148- linear_model = linear .LinearModel (units = 1 )
149- dnn_model = sequential .Sequential ([core .Dense (units = 1 )])
150- wide_deep_model = wide_deep .WideDeepModel (linear_model , dnn_model )
151- linear_opt = gradient_descent .SGD (learning_rate = 0.05 )
152- dnn_opt = adagrad .Adagrad (learning_rate = 0.1 )
153- wide_deep_model .compile (optimizer = [linear_opt , dnn_opt ], loss = "mse" )
154-
155- if use_dataset_creator :
156- x = dataset_creator .DatasetCreator (dataset_fn )
157- hist = wide_deep_model .fit (
147+ with distribution .scope ():
148+ linear_model = linear .LinearModel (units = 1 )
149+ dnn_model = sequential .Sequential ([core .Dense (units = 1 )])
150+ wide_deep_model = wide_deep .WideDeepModel (linear_model , dnn_model )
151+ linear_opt = gradient_descent .SGD (learning_rate = 0.05 )
152+ dnn_opt = adagrad .Adagrad (learning_rate = 0.1 )
153+ wide_deep_model .compile (optimizer = [linear_opt , dnn_opt ], loss = "mse" )
154+
155+ if use_dataset_creator :
156+ x = dataset_creator .DatasetCreator (dataset_fn )
157+ hist = wide_deep_model .fit (
158158 x , epochs = 3 , steps_per_epoch = INPUT_SIZE
159159 )
160- else :
161- if data_fn == "numpy" :
162- inputs , output = get_numpy ()
163- hist = wide_deep_model .fit (inputs , output , epochs = 3 )
164- else :
165- hist = wide_deep_model .fit (get_dataset (), epochs = 3 )
166- self .assertLess (hist .history ["loss" ][2 ], 0.2 )
160+ else :
161+ if data_fn == "numpy" :
162+ inputs , output = get_numpy ()
163+ hist = wide_deep_model .fit (inputs , output , epochs = 3 )
164+ else :
165+ hist = wide_deep_model .fit (
166+ get_dataset (), epochs = 3 , steps_per_epoch = INPUT_SIZE
167+ )
168+ self .assertLess (hist .history ["loss" ][2 ], 0.2 )
167169
168170
169171if __name__ == "__main__" :
0 commit comments