FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

word_cnn training completed, but saving model doesn't work. · MSavameri/TensorFlow.NET@e56e5d3 · GitHub

Commit e56e5d3

Browse files
committed
word_cnn training completed, but saving model doesn't work.
1 parent 7a706c9 commit e56e5d3

6 files changed

Lines changed: 90 additions & 28 deletions

File tree

‎data/dbpedia_subset.zip‎

39.7 KB
Binary file not shown.

‎graph/word_cnn.meta‎

85.5 KB
Binary file not shown.

‎src/TensorFlowNET.Core/APIs/tf.variable.cs‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,12 @@ namespace Tensorflow
66
{
77
public static partial class tf
88
{
9+
public static VariableV1[] global_variables(string scope = null)
10+
{
11+
return (ops.get_collection(ops.GraphKeys.GLOBAL_VARIABLES, scope) as List<VariableV1>)
12+
.ToArray();
13+
}
14+
915
public static Operation global_variables_initializer()
1016
{
1117
var g = variables.global_variables();

‎src/TensorFlowNET.Core/Train/tf.optimizers.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@ public static class train
1414

1515
public static Optimizer AdamOptimizer(float learning_rate) => new AdamOptimizer(learning_rate);
1616

17-
public static Saver Saver() => new Saver();
17+
public static Saver Saver(VariableV1[] var_list = null) => new Saver(var_list: var_list);
1818

1919
public static string write_graph(Graph graph, string logdir, string name, bool as_text = true)
2020
=> graph_io.write_graph(graph, logdir, name, as_text);

‎test/TensorFlowNET.Examples/TextProcess/DataHelpers.cs‎

Lines changed: 42 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,46 @@ namespace TensorFlowNET.Examples
1212
{
1313
public class DataHelpers
1414
{
15+
public static Dictionary<string, int> build_word_dict(string path)
16+
{
17+
var contents = File.ReadAllLines(path);
18+
19+
var words = new List<string>();
20+
foreach (var content in contents)
21+
words.AddRange(clean_str(content).Split(' ').Where(x => x.Length > 1));
22+
var word_counter = words.GroupBy(x => x)
23+
.Select(x => new { Word = x.Key, Count = x.Count() })
24+
.OrderByDescending(x => x.Count)
25+
.ToArray();
26+
27+
var word_dict = new Dictionary<string, int>();
28+
word_dict["<pad>"] = 0;
29+
word_dict["<unk>"] = 1;
30+
word_dict["<eos>"] = 2;
31+
foreach (var word in word_counter)
32+
word_dict[word.Word] = word_dict.Count;
33+
34+
return word_dict;
35+
}
36+
37+
public static (int[][], int[]) build_word_dataset(string path, Dictionary<string, int> word_dict, int document_max_len)
38+
{
39+
var contents = File.ReadAllLines(path);
40+
var x = contents.Select(c => (clean_str(c) + " <eos>")
41+
.Split(' ').Take(document_max_len)
42+
.Select(w => word_dict.ContainsKey(w) ? word_dict[w] : word_dict["<unk>"]).ToArray())
43+
.ToArray();
44+
45+
for (int i = 0; i < x.Length; i++)
46+
if (x[i].Length == document_max_len)
47+
x[i][document_max_len - 1] = word_dict["<eos>"];
48+
else
49+
Array.Resize(ref x[i], document_max_len);
50+
51+
var y = contents.Select(c => int.Parse(c.Substring(0, c.IndexOf(','))) - 1).ToArray();
52+
53+
return (x, y);
54+
}
1555

1656
public static (int[][], int[], int) build_char_dataset(string path, string model, int document_max_len, int? limit = null, bool shuffle=true)
1757
{
@@ -96,8 +136,8 @@ public static (string[], NDArray) load_data_and_labels(string positive_data_file
96136

97137
private static string clean_str(string str)
98138
{
99-
str = Regex.Replace(str, @"[^A-Za-z0-9(),!?\'\`]", " ");
100-
str = Regex.Replace(str, @"\'s", " \'s");
139+
str = Regex.Replace(str, "[^A-Za-z0-9(),!?]", " ");
140+
str = Regex.Replace(str, ",", " ");
101141
return str;
102142
}
103143

‎test/TensorFlowNET.Examples/TextProcess/TextClassificationTrain.cs‎

Lines changed: 41 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -26,22 +26,23 @@ public class TextClassificationTrain : IExample
2626
public string Name => "Text Classification";
2727
public int? DataLimit = null;
2828
public bool ImportGraph { get; set; } = true;
29-
public bool UseSubset = true; // <----- set this true to use a limited subset of dbpedia
29+
public bool UseSubset = false; // <----- set this true to use a limited subset of dbpedia
3030

3131
private string dataDir = "text_classification";
3232
private string dataFileName = "dbpedia_csv.tar.gz";
3333

34-
public string model_name = "vd_cnn"; // word_cnn | char_cnn | vd_cnn | word_rnn | att_rnn | rcnn
34+
public string model_name = "word_cnn"; // word_cnn | char_cnn | vd_cnn | word_rnn | att_rnn | rcnn
3535

3636
private const string TRAIN_PATH = "text_classification/dbpedia_csv/train.csv";
3737
private const string SUBSET_PATH = "text_classification/dbpedia_csv/dbpedia_6400.csv";
3838
private const string TEST_PATH = "text_classification/dbpedia_csv/test.csv";
3939

40-
private const int CHAR_MAX_LEN = 1014;
41-
private const int WORD_MAX_LEN = 1014;
4240
private const int NUM_CLASS = 14;
4341
private const int BATCH_SIZE = 64;
4442
private const int NUM_EPOCHS = 10;
43+
private const int WORD_MAX_LEN = 100;
44+
private const int CHAR_MAX_LEN = 1014;
45+
4546
protected float loss_value = 0;
4647

4748
public bool Run()
@@ -61,8 +62,21 @@ protected virtual bool RunWithImportedGraph(Session sess, Graph graph)
6162
{
6263
var stopwatch = Stopwatch.StartNew();
6364
Console.WriteLine("Building dataset...");
64-
var path = UseSubset ? SUBSET_PATH : TRAIN_PATH;
65-
var (x, y, alphabet_size) = DataHelpers.build_char_dataset(path, model_name, CHAR_MAX_LEN, DataLimit = null, shuffle:!UseSubset);
65+
var path = UseSubset ? SUBSET_PATH : TRAIN_PATH;
66+
int[][] x = null;
67+
int[] y = null;
68+
int alphabet_size = 0;
69+
int vocabulary_size = 0;
70+
71+
if (model_name == "vd_cnn")
72+
(x, y, alphabet_size) = DataHelpers.build_char_dataset(path, model_name, CHAR_MAX_LEN, DataLimit = null, shuffle:!UseSubset);
73+
else
74+
{
75+
var word_dict = DataHelpers.build_word_dict(TRAIN_PATH);
76+
vocabulary_size = len(word_dict);
77+
(x, y) = DataHelpers.build_word_dataset(TRAIN_PATH, word_dict, WORD_MAX_LEN);
78+
}
79+
6680
Console.WriteLine("\tDONE ");
6781

6882
var (train_x, valid_x, train_y, valid_y) = train_test_split(x, y, test_size: 0.15f);
@@ -75,18 +89,19 @@ protected virtual bool RunWithImportedGraph(Session sess, Graph graph)
7589
Console.WriteLine("\tDONE " + stopwatch.Elapsed);
7690

7791
sess.run(tf.global_variables_initializer());
92+
var saver = tf.train.Saver(tf.global_variables());
7893

7994
var train_batches = batch_iter(train_x, train_y, BATCH_SIZE, NUM_EPOCHS);
8095
var num_batches_per_epoch = (len(train_x) - 1) / BATCH_SIZE + 1;
8196
double max_accuracy = 0;
8297

83-
Tensor is_training = graph.get_tensor_by_name("is_training:0");
84-
Tensor model_x = graph.get_tensor_by_name("x:0");
85-
Tensor model_y = graph.get_tensor_by_name("y:0");
86-
Tensor loss = graph.get_tensor_by_name("loss/value:0");
87-
Tensor optimizer = graph.get_tensor_by_name("loss/optimizer:0");
88-
Tensor global_step = graph.get_tensor_by_name("global_step:0");
89-
Tensor accuracy = graph.get_tensor_by_name("accuracy/value:0");
98+
Tensor is_training = graph.OperationByName("is_training");
99+
Tensor model_x = graph.OperationByName("x");
100+
Tensor model_y = graph.OperationByName("y");
101+
Tensor loss = graph.OperationByName("loss/Mean"); // word_cnn
102+
Operation optimizer = graph.OperationByName("loss/Adam"); // word_cnn
103+
Tensor global_step = graph.OperationByName("Variable");
104+
Tensor accuracy = graph.OperationByName("accuracy/accuracy");
90105
stopwatch = Stopwatch.StartNew();
91106
int i = 0;
92107
foreach (var (x_batch, y_batch, total) in train_batches)
@@ -105,11 +120,10 @@ protected virtual bool RunWithImportedGraph(Session sess, Graph graph)
105120
var result = sess.run(new ITensorOrOperation[] { optimizer, global_step, loss }, train_feed_dict);
106121
loss_value = result[2];
107122
var step = (int)result[1];
108-
if (step % 10 == 0 || step < 10)
123+
if (step % 10 == 0)
109124
{
110125
var estimate = TimeSpan.FromSeconds((stopwatch.Elapsed.TotalSeconds / i) * total);
111-
Console.WriteLine($"Training on batch {i}/{total}. Estimated training time: {estimate}");
112-
Console.WriteLine($"Step {step} loss: {loss_value}");
126+
Console.WriteLine($"Training on batch {i}/{total} loss: {loss_value}. Estimated training time: {estimate}");
113127
}
114128

115129
if (step % 100 == 0)
@@ -133,13 +147,15 @@ protected virtual bool RunWithImportedGraph(Session sess, Graph graph)
133147

134148
var valid_accuracy = sum_accuracy / cnt;
135149

136-
print($"\nValidation Accuracy = {valid_accuracy}\n");
137-
138-
// # Save model
139-
// if valid_accuracy > max_accuracy:
140-
// max_accuracy = valid_accuracy
141-
// saver.save(sess, "{0}/{1}.ckpt".format(args.model, args.model), global_step = step)
142-
// print("Model is saved.\n")
150+
print($"\nValidation Accuracy = {valid_accuracy}\n");
151+
152+
// # Save model
153+
if (valid_accuracy > max_accuracy)
154+
{
155+
max_accuracy = valid_accuracy;
156+
// saver.save(sess, $"{dataDir}/{model_name}.ckpt", global_step: step.ToString());
157+
print("Model is saved.\n");
158+
}
143159
}
144160
}
145161

@@ -180,9 +196,9 @@ protected virtual bool RunWithBuiltGraph(Session session, Graph graph)
180196
//int samples = len / classes;
181197
int train_size = (int)Math.Round(len * (1 - test_size));
182198
var train_x = x[new Slice(stop: train_size), new Slice()];
183-
var valid_x = x[new Slice(start: train_size + 1), new Slice()];
199+
var valid_x = x[new Slice(start: train_size), new Slice()];
184200
var train_y = y[new Slice(stop: train_size)];
185-
var valid_y = y[new Slice(start: train_size + 1)];
201+
var valid_y = y[new Slice(start: train_size)];
186202
Console.WriteLine("\tDONE");
187203
return (train_x, valid_x, train_y, valid_y);
188204
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL