| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 7a706c9 commit e56e5d3
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -6,6 +6,12 @@ namespace Tensorflow | |||
| 6 | 6 | { | |
| 7 | 7 | public static partial class tf | |
| 8 | 8 | { | |
| 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 | + | ||
| 9 | 15 | public static Operation global_variables_initializer() | |
| 10 | 16 | { | |
| 11 | 17 | var g = variables.global_variables(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,7 +14,7 @@ public static class train | |||
| 14 | 14 | ||
| 15 | 15 | public static Optimizer AdamOptimizer(float learning_rate) => new AdamOptimizer(learning_rate); | |
| 16 | 16 | ||
| 17 | - public static Saver Saver() => new Saver(); | ||
| 17 | + public static Saver Saver(VariableV1[] var_list = null) => new Saver(var_list: var_list); | ||
| 18 | 18 | ||
| 19 | 19 | public static string write_graph(Graph graph, string logdir, string name, bool as_text = true) | |
| 20 | 20 | => graph_io.write_graph(graph, logdir, name, as_text); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,6 +12,46 @@ namespace TensorFlowNET.Examples | |||
| 12 | 12 | { | |
| 13 | 13 | public class DataHelpers | |
| 14 | 14 | { | |
| 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 | + } | ||
| 15 | 55 | ||
| 16 | 56 | public static (int[][], int[], int) build_char_dataset(string path, string model, int document_max_len, int? limit = null, bool shuffle=true) | |
| 17 | 57 | { | |
@@ -96,8 +136,8 @@ public static (string[], NDArray) load_data_and_labels(string positive_data_file | |||
| 96 | 136 | ||
| 97 | 137 | private static string clean_str(string str) | |
| 98 | 138 | { | |
| 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, ",", " "); | ||
| 101 | 141 | return str; | |
| 102 | 142 | } | |
| 103 | 143 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -26,22 +26,23 @@ public class TextClassificationTrain : IExample | |||
| 26 | 26 | public string Name => "Text Classification"; | |
| 27 | 27 | public int? DataLimit = null; | |
| 28 | 28 | 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 | ||
| 30 | 30 | ||
| 31 | 31 | private string dataDir = "text_classification"; | |
| 32 | 32 | private string dataFileName = "dbpedia_csv.tar.gz"; | |
| 33 | 33 | ||
| 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 | ||
| 35 | 35 | ||
| 36 | 36 | private const string TRAIN_PATH = "text_classification/dbpedia_csv/train.csv"; | |
| 37 | 37 | private const string SUBSET_PATH = "text_classification/dbpedia_csv/dbpedia_6400.csv"; | |
| 38 | 38 | private const string TEST_PATH = "text_classification/dbpedia_csv/test.csv"; | |
| 39 | 39 | ||
| 40 | - private const int CHAR_MAX_LEN = 1014; | ||
| 41 | - private const int WORD_MAX_LEN = 1014; | ||
| 42 | 40 | private const int NUM_CLASS = 14; | |
| 43 | 41 | private const int BATCH_SIZE = 64; | |
| 44 | 42 | private const int NUM_EPOCHS = 10; | |
| 43 | + private const int WORD_MAX_LEN = 100; | ||
| 44 | + private const int CHAR_MAX_LEN = 1014; | ||
| 45 | + | ||
| 45 | 46 | protected float loss_value = 0; | |
| 46 | 47 | ||
| 47 | 48 | public bool Run() | |
@@ -61,8 +62,21 @@ protected virtual bool RunWithImportedGraph(Session sess, Graph graph) | |||
| 61 | 62 | { | |
| 62 | 63 | var stopwatch = Stopwatch.StartNew(); | |
| 63 | 64 | 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 | + | ||
| 66 | 80 | Console.WriteLine("\tDONE "); | |
| 67 | 81 | ||
| 68 | 82 | 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) | |||
| 75 | 89 | Console.WriteLine("\tDONE " + stopwatch.Elapsed); | |
| 76 | 90 | ||
| 77 | 91 | sess.run(tf.global_variables_initializer()); | |
| 92 | + var saver = tf.train.Saver(tf.global_variables()); | ||
| 78 | 93 | ||
| 79 | 94 | var train_batches = batch_iter(train_x, train_y, BATCH_SIZE, NUM_EPOCHS); | |
| 80 | 95 | var num_batches_per_epoch = (len(train_x) - 1) / BATCH_SIZE + 1; | |
| 81 | 96 | double max_accuracy = 0; | |
| 82 | 97 | ||
| 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"); | ||
| 90 | 105 | stopwatch = Stopwatch.StartNew(); | |
| 91 | 106 | int i = 0; | |
| 92 | 107 | foreach (var (x_batch, y_batch, total) in train_batches) | |
@@ -105,11 +120,10 @@ protected virtual bool RunWithImportedGraph(Session sess, Graph graph) | |||
| 105 | 120 | var result = sess.run(new ITensorOrOperation[] { optimizer, global_step, loss }, train_feed_dict); | |
| 106 | 121 | loss_value = result[2]; | |
| 107 | 122 | var step = (int)result[1]; | |
| 108 | - if (step % 10 == 0 || step < 10) | ||
| 123 | + if (step % 10 == 0) | ||
| 109 | 124 | { | |
| 110 | 125 | 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}"); | ||
| 113 | 127 | } | |
| 114 | 128 | ||
| 115 | 129 | if (step % 100 == 0) | |
@@ -133,13 +147,15 @@ protected virtual bool RunWithImportedGraph(Session sess, Graph graph) | |||
| 133 | 147 | ||
| 134 | 148 | var valid_accuracy = sum_accuracy / cnt; | |
| 135 | 149 | ||
| 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 | + } | ||
| 143 | 159 | } | |
| 144 | 160 | } | |
| 145 | 161 | ||
@@ -180,9 +196,9 @@ protected virtual bool RunWithBuiltGraph(Session session, Graph graph) | |||
| 180 | 196 | //int samples = len / classes; | |
| 181 | 197 | int train_size = (int)Math.Round(len * (1 - test_size)); | |
| 182 | 198 | 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()]; | ||
| 184 | 200 | 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)]; | ||
| 186 | 202 | Console.WriteLine("\tDONE"); | |
| 187 | 203 | return (train_x, valid_x, train_y, valid_y); | |
| 188 | 204 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments