| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 4c1878b commit 6a9ccea
15 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -25,5 +25,32 @@ public struct TF_Buffer | |||
| 25 | 25 | public IntPtr data; | |
| 26 | 26 | public ulong length; | |
| 27 | 27 | public IntPtr data_deallocator; | |
| 28 | + | ||
| 29 | + public unsafe Span<T> AsSpan<T>() where T: unmanaged | ||
| 30 | + { | ||
| 31 | + if(length > int.MaxValue) | ||
| 32 | + { | ||
| 33 | + throw new ValueError($"The length {length} is too large to use in the span."); | ||
| 34 | + } | ||
| 35 | + return new Span<T>(data.ToPointer(), (int)length); | ||
| 36 | + } | ||
| 37 | + | ||
| 38 | + public unsafe byte[] ToByteArray() | ||
| 39 | + { | ||
| 40 | + byte[] res = new byte[length]; | ||
| 41 | + if(length > int.MaxValue) | ||
| 42 | + { | ||
| 43 | + byte* root = (byte*)data; | ||
| 44 | + for(ulong i = 0; i < length; i++) | ||
| 45 | + { | ||
| 46 | + res[i] = *(root++); | ||
| 47 | + } | ||
| 48 | + } | ||
| 49 | + else | ||
| 50 | + { | ||
| 51 | + new Span<byte>(data.ToPointer(), (int)length).CopyTo(res.AsSpan()); | ||
| 52 | + } | ||
| 53 | + return res; | ||
| 54 | + } | ||
| 28 | 55 | } | |
| 29 | 56 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -18,6 +18,10 @@ public static (DataType[], Tensor[]) onvert_to_mixed_eager_tensors(Tensor[] valu | |||
| 18 | 18 | var types = v.Select(t => t.dtype.as_datatype_enum()); | |
| 19 | 19 | return (types.ToArray(), v.ToArray()); | |
| 20 | 20 | } | |
| 21 | + public static Tensor[] executes(string op_name, int num_outputs, Tensor[] inputs, object[] attrs, Context ctx, string name = null) | ||
| 22 | + { | ||
| 23 | + return quick_execute(op_name, num_outputs, inputs, attrs, ctx, name); | ||
| 24 | + } | ||
| 21 | 25 | public static Tensor[] quick_execute(string op_name, int num_outputs, Tensor[] inputs, object[] attrs, Context ctx, string name = null) | |
| 22 | 26 | { | |
| 23 | 27 | string device_name = ctx.DeviceName; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -149,6 +149,7 @@ private static void _ProcessNewOps(Graph graph) | |||
| 149 | 149 | foreach (var new_op in graph._add_new_tf_operations()) | |
| 150 | 150 | { | |
| 151 | 151 | var original_device = new_op.Device; | |
| 152 | + new_op._set_device(original_device); | ||
| 152 | 153 | } | |
| 153 | 154 | } | |
| 154 | 155 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,9 +1,11 @@ | |||
| 1 | 1 | using Google.Protobuf; | |
| 2 | 2 | using System; | |
| 3 | 3 | using System.Collections.Generic; | |
| 4 | + using System.IO; | ||
| 4 | 5 | using System.Linq; | |
| 5 | 6 | using System.Text; | |
| 6 | 7 | using Tensorflow.Contexts; | |
| 8 | + using Tensorflow.Eager; | ||
| 7 | 9 | using Tensorflow.Graphs; | |
| 8 | 10 | using Tensorflow.Operations; | |
| 9 | 11 | using Tensorflow.Util; | |
@@ -16,6 +18,8 @@ public class EagerDefinedFunction | |||
| 16 | 18 | public int _num_outputs; | |
| 17 | 19 | FuncGraph _func_graph; | |
| 18 | 20 | FunctionDef _definition; | |
| 21 | + OpDef _signature; | ||
| 22 | + string _name; | ||
| 19 | 23 | Tensor[] _func_graph_outputs; | |
| 20 | 24 | public string Name => _func_graph.FuncName; | |
| 21 | 25 | public DataType[] OutputTypes { get; protected set; } | |
@@ -31,6 +35,18 @@ public FunctionDef Definition | |||
| 31 | 35 | return _definition; | |
| 32 | 36 | } | |
| 33 | 37 | } | |
| 38 | + | ||
| 39 | + public OpDef Signature | ||
| 40 | + { | ||
| 41 | + get | ||
| 42 | + { | ||
| 43 | + if( _signature is null) | ||
| 44 | + { | ||
| 45 | + _signature = Definition.Signature; | ||
| 46 | + } | ||
| 47 | + return _signature; | ||
| 48 | + } | ||
| 49 | + } | ||
| 34 | 50 | public EagerDefinedFunction(string name, FuncGraph graph, | |
| 35 | 51 | Tensors inputs, Tensors outputs, | |
| 36 | 52 | Dictionary<string, string> attrs) | |
@@ -75,12 +91,12 @@ public Tensors Call(Tensors args) | |||
| 75 | 91 | Tensor[] outputs; | |
| 76 | 92 | if (executing_eagerly) | |
| 77 | 93 | { | |
| 78 | - outputs = tf.Runner.TFE_Execute(tf.Context, | ||
| 79 | - tf.Context.DeviceName, | ||
| 80 | - _func_graph.FuncName, | ||
| 81 | - args, | ||
| 82 | - attrs, | ||
| 83 | - _num_outputs); | ||
| 94 | + outputs = execute.executes( | ||
| 95 | + Signature.Name, | ||
| 96 | + _num_outputs, | ||
| 97 | + args, | ||
| 98 | + attrs, | ||
| 99 | + tf.Context); | ||
| 84 | 100 | } | |
| 85 | 101 | else | |
| 86 | 102 | { | |
@@ -135,9 +151,13 @@ public void AddToGraph(Graph g = null) | |||
| 135 | 151 | private FunctionDef _get_definition() | |
| 136 | 152 | { | |
| 137 | 153 | var buffer = c_api_util.tf_buffer(); | |
| 138 | - // TODO(Rinne): pywrap_tf_session.TF_FunctionToFunctionDef | ||
| 154 | + Status status = new(); | ||
| 155 | + c_api.TF_FunctionToFunctionDef(_func_graph._func_graph_handle, buffer, status); | ||
| 156 | + status.Check(true); | ||
| 139 | 157 | var proto_data = c_api.TF_GetBuffer(buffer); | |
| 140 | - throw new NotImplementedException(); | ||
| 158 | + FunctionDef function_def = new(); | ||
| 159 | + function_def.MergeFrom(proto_data.AsSpan<byte>()); | ||
| 160 | + return function_def; | ||
| 141 | 161 | } | |
| 142 | 162 | } | |
| 143 | 163 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,7 +10,7 @@ namespace Tensorflow.Graphs; | |||
| 10 | 10 | /// </summary> | |
| 11 | 11 | public class FuncGraph : Graph, IDisposable | |
| 12 | 12 | { | |
| 13 | - SafeFuncGraphHandle _func_graph_handle; | ||
| 13 | + internal SafeFuncGraphHandle _func_graph_handle; | ||
| 14 | 14 | public string FuncName => _graph_key; | |
| 15 | 15 | ||
| 16 | 16 | public Tensors Inputs { get; set; } = new Tensors(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -238,6 +238,19 @@ public TF_AttrMetadata GetAttributeMetadata(string attr_name, Status s) | |||
| 238 | 238 | return c_api.TF_OperationGetAttrMetadata(_handle, attr_name, s); | |
| 239 | 239 | } | |
| 240 | 240 | ||
| 241 | + [Obsolete("The implementation is not complete.")] | ||
| 242 | + internal void _set_device_from_string(string device_str) | ||
| 243 | + { | ||
| 244 | + // TODO(Rinne): complete it with new C API `SetRequestedDevice`. | ||
| 245 | + //c_api.TF_SetDevice(_handle, device_str); | ||
| 246 | + } | ||
| 247 | + | ||
| 248 | + [Obsolete("The implementation is not complete.")] | ||
| 249 | + internal void _set_device(string device) | ||
| 250 | + { | ||
| 251 | + _set_device_from_string(device); | ||
| 252 | + } | ||
| 253 | + | ||
| 241 | 254 | private NodeDef GetNodeDef() | |
| 242 | 255 | { | |
| 243 | 256 | var buffer = new Buffer(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -45,11 +45,8 @@ public Loader(SavedObjectGraph object_graph_proto, SavedModel saved_model_proto, | |||
| 45 | 45 | _asset_file_def = meta_graph.AssetFileDef; | |
| 46 | 46 | _operation_attributes = meta_graph.GraphDef.Node.ToDictionary(x => x.Name, x => x.Attr); | |
| 47 | 47 | _proto = object_graph_proto; | |
| 48 | - // Debug(Rinne) | ||
| 49 | - var temp = _proto.ToString(); | ||
| 50 | 48 | _export_dir = export_dir; | |
| 51 | - // TODO: `this._concrete_functions` and `this._restored_concrete_functions` | ||
| 52 | - // TODO(Rinne): This method is very slow, needs to be accelareted. | ||
| 49 | + // TODO(Rinne): This method is a bit slow (especially under debug mode), may need to be accelareted. | ||
| 53 | 50 | _concrete_functions = function_deserialization.load_function_def_library( | |
| 54 | 51 | meta_graph.GraphDef.Library, _proto); | |
| 55 | 52 | _restored_concrete_functions = new HashSet<string>(); | |
@@ -322,11 +319,6 @@ private void _load_checkpoint_save_and_restore_functions() | |||
| 322 | 319 | foreach(var (node_id, proto) in _iter_all_nodes()) | |
| 323 | 320 | { | |
| 324 | 321 | var node = get(node_id); | |
| 325 | - if(node is null) | ||
| 326 | - { | ||
| 327 | - // skip it because now we skip the restoration of `Function` and `ConcreteFunction`. | ||
| 328 | - continue; | ||
| 329 | - } | ||
| 330 | 322 | if(proto.SaveableObjects.Keys.Count == 1 && proto.SaveableObjects.First().Key == TrackableUtils.SERIALIZE_TO_TENSORS_NAME) | |
| 331 | 323 | { | |
| 332 | 324 | // Restore Trackable serialize- and restore-from-tensor functions. | |
@@ -390,7 +382,7 @@ private void _load_nodes() | |||
| 390 | 382 | var optimizer_object = nodes[optimizer_node_id]; | |
| 391 | 383 | var optimizer_variable = nodes[slot_variable_proto.OriginalVariableNodeId]; | |
| 392 | 384 | ||
| 393 | - // TODO: implement it. | ||
| 385 | + // TODO(Rinne): implement it. | ||
| 394 | 386 | throw new NotImplementedException("The model loading of SavedModel still has some incompleted part." + | |
| 395 | 387 | " Please submit an issue to https://github.com/SciSharp/TensorFlow.NET/issues."); | |
| 396 | 388 | } | |
@@ -508,21 +500,11 @@ public Trackable get(string node_id) | |||
| 508 | 500 | /// <param name="node_id"></param> | |
| 509 | 501 | private void _add_object_graph_edges(SavedObject proto, int node_id) | |
| 510 | 502 | { | |
| 511 | - // Debug(Rinne) | ||
| 512 | - if(node_id == 1) | ||
| 513 | - { | ||
| 514 | - Console.WriteLine(); | ||
| 515 | - } | ||
| 516 | 503 | var obj = _nodes[node_id]; | |
| 517 | 504 | var setter = _node_setters[node_id]; | |
| 518 | 505 | ||
| 519 | 506 | foreach(var refer in proto.Children) | |
| 520 | 507 | { | |
| 521 | - if(obj is null) | ||
| 522 | - { | ||
| 523 | - // skip it because now we skip the restoration of `Function` and `ConcreteFunction`. | ||
| 524 | - continue; | ||
| 525 | - } | ||
| 526 | 508 | setter.Invoke(obj, refer.LocalName, _nodes[refer.NodeId]); | |
| 527 | 509 | // TODO(Rinne): deal with "__call__" | |
| 528 | 510 | } | |
@@ -553,12 +535,6 @@ private void _add_object_graph_edges(SavedObject proto, int node_id) | |||
| 553 | 535 | private (Trackable, Action<object, object, object>) _recreate(SavedObject proto, int node_id, IDictionary<int, Trackable> nodes) | |
| 554 | 536 | { | |
| 555 | 537 | // skip the registered classes. | |
| 556 | - if(node_id == 16) | ||
| 557 | - { | ||
| 558 | - // Debug(Rinne) | ||
| 559 | - Console.WriteLine(); | ||
| 560 | - } | ||
| 561 | - | ||
| 562 | 538 | Dictionary<OneOf<string, int>, Trackable> dependencies = new(); | |
| 563 | 539 | foreach(var item in _get_node_dependencies(proto)) | |
| 564 | 540 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -65,6 +65,8 @@ public BaseResourceVariable() | |||
| 65 | 65 | } | |
| 66 | 66 | ||
| 67 | 67 | public void __init__(bool trainable = true, | |
| 68 | + Shape shape = null, | ||
| 69 | + TF_DataType dtype = TF_DataType.DtInvalid, | ||
| 68 | 70 | Tensor handle = null, | |
| 69 | 71 | string name = null, | |
| 70 | 72 | string unique_id = null, | |
@@ -75,6 +77,14 @@ public void __init__(bool trainable = true, | |||
| 75 | 77 | _unique_id = unique_id; | |
| 76 | 78 | this.handle = handle; | |
| 77 | 79 | _name = name; | |
| 80 | + if(shape is not null) | ||
| 81 | + { | ||
| 82 | + _shape = shape; | ||
| 83 | + } | ||
| 84 | + if(dtype != TF_DataType.DtInvalid) | ||
| 85 | + { | ||
| 86 | + _dtype = dtype; | ||
| 87 | + } | ||
| 78 | 88 | ||
| 79 | 89 | // After the handle has been created, set up a way to clean it up when | |
| 80 | 90 | // executing eagerly. We'll hold the only reference to the deleter, so that | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -116,7 +116,11 @@ private void _init_from_args(object initial_value = null, | |||
| 116 | 116 | } | |
| 117 | 117 | }); | |
| 118 | 118 | ||
| 119 | - _shape = shape ?? _initial_value.shape; | ||
| 119 | + if(shape is null) | ||
| 120 | + { | ||
| 121 | + shape = _initial_value.shape; | ||
| 122 | + } | ||
| 123 | + dtype = _initial_value.dtype; | ||
| 120 | 124 | ||
| 121 | 125 | if (_in_graph_mode) | |
| 122 | 126 | { | |
@@ -135,7 +139,7 @@ private void _init_from_args(object initial_value = null, | |||
| 135 | 139 | { | |
| 136 | 140 | handle = resource_variable_ops.eager_safe_variable_handle( | |
| 137 | 141 | initial_value: _initial_value, | |
| 138 | - shape: _shape, | ||
| 142 | + shape: shape, | ||
| 139 | 143 | shared_name: shared_name, | |
| 140 | 144 | name: name, | |
| 141 | 145 | graph_mode: _in_graph_mode); | |
@@ -154,6 +158,8 @@ private void _init_from_args(object initial_value = null, | |||
| 154 | 158 | } | |
| 155 | 159 | ||
| 156 | 160 | base.__init__(trainable: trainable, | |
| 161 | + shape: shape, | ||
| 162 | + dtype: dtype, | ||
| 157 | 163 | handle: handle, | |
| 158 | 164 | name: name, | |
| 159 | 165 | unique_id: unique_id, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -50,9 +50,9 @@ public UninitializedVariable( | |||
| 50 | 50 | { | |
| 51 | 51 | tf_with(ops.name_scope("Read"), _ => | |
| 52 | 52 | { | |
| 53 | - tf.device(handle.Device); | ||
| 54 | - var value = gen_resource_variable_ops.read_variable_op(handle, dtype); | ||
| 55 | - resource_variable_ops._maybe_set_handle_data(dtype, handle, value); | ||
| 53 | + tf.device(created_handle.Device); | ||
| 54 | + var value = gen_resource_variable_ops.read_variable_op(created_handle, dtype); | ||
| 55 | + resource_variable_ops._maybe_set_handle_data(dtype, created_handle, value); | ||
| 56 | 56 | _graph_element = value; | |
| 57 | 57 | }); | |
| 58 | 58 | ops.add_to_collection(ops.GraphKeys.GLOBAL_VARIABLES_, this); | |
@@ -63,9 +63,7 @@ public UninitializedVariable( | |||
| 63 | 63 | } | |
| 64 | 64 | }); | |
| 65 | 65 | }); | |
| 66 | - _shape = shape; | ||
| 67 | - _dtype = dtype; | ||
| 68 | - base.__init__(trainable, created_handle, unique_id: unique_id, handle_name: handle_name); | ||
| 66 | + base.__init__(trainable, shape, dtype, created_handle, unique_id: unique_id, handle_name: handle_name); | ||
| 69 | 67 | } | |
| 70 | 68 | } | |
| 71 | 69 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments