| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent e93f112 commit 7764865
13 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -9,7 +9,7 @@ | |||
| 9 | 9 | [](https://996.icu/#/en_US) | |
| 10 | 10 | [](https://mybinder.org/v2/gh/javiercp/BinderTF.NET/master?urlpath=lab) | |
| 11 | 11 | ||
| 12 | - *master branch is based on tensorflow 2.1 now, v0.15-tensorflow1.15 is from tensorflow1.15.* | ||
| 12 | + *master branch is based on tensorflow 2.2 now, v0.15-tensorflow1.15 is from tensorflow1.15.* | ||
| 13 | 13 | ||
| 14 | 14 | TF.NET is a member project of [SciSharp STACK](https://github.com/SciSharp). | |
| 15 | 15 | ||
@@ -28,7 +28,7 @@ In comparison to other projects, like for instance TensorFlowSharp which only pr | |||
| 28 | 28 | ||
| 29 | 29 | ### How to use | |
| 30 | 30 | ||
| 31 | - | TensorFlow | tf 1.13 | tf 1.14 | tf 1.15 | tf 2.0 | | ||
| 31 | + | TensorFlow | tf 1.13 | tf 1.14 | tf 1.15 | tf 2.2 | | ||
| 32 | 32 | | ----------- | ------- | ------- | ------- | ------ | | |
| 33 | 33 | | tf.net 0.20 | | | x | x | | |
| 34 | 34 | | tf.net 0.15 | | x | x | | | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -13,31 +13,37 @@ public EagerTensor(IntPtr handle) : base(handle) | |||
| 13 | 13 | { | |
| 14 | 14 | tfe_tensor_handle = handle; | |
| 15 | 15 | _handle = c_api.TFE_TensorHandleResolve(handle, status); | |
| 16 | + _id = ops.uid(); | ||
| 16 | 17 | } | |
| 17 | 18 | ||
| 18 | 19 | public EagerTensor(string value, string device_name) : base(value) | |
| 19 | 20 | { | |
| 20 | 21 | tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | |
| 22 | + _id = ops.uid(); | ||
| 21 | 23 | } | |
| 22 | 24 | ||
| 23 | 25 | public EagerTensor(int value, string device_name) : base(value) | |
| 24 | 26 | { | |
| 25 | 27 | tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | |
| 28 | + _id = ops.uid(); | ||
| 26 | 29 | } | |
| 27 | 30 | ||
| 28 | 31 | public EagerTensor(float[] value, string device_name) : base(value) | |
| 29 | 32 | { | |
| 30 | 33 | tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | |
| 34 | + _id = ops.uid(); | ||
| 31 | 35 | } | |
| 32 | 36 | ||
| 33 | 37 | public EagerTensor(double[] value, string device_name) : base(value) | |
| 34 | 38 | { | |
| 35 | 39 | tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | |
| 40 | + _id = ops.uid(); | ||
| 36 | 41 | } | |
| 37 | 42 | ||
| 38 | 43 | public EagerTensor(NDArray value, string device_name) : base(value) | |
| 39 | 44 | { | |
| 40 | 45 | tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | |
| 46 | + _id = ops.uid(); | ||
| 41 | 47 | } | |
| 42 | 48 | ||
| 43 | 49 | public override string ToString() | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -102,14 +102,20 @@ public partial class c_api | |||
| 102 | 102 | public static extern TFE_Op TFE_NewOp(IntPtr ctx, string op_or_function_name, IntPtr status); | |
| 103 | 103 | ||
| 104 | 104 | /// <summary> | |
| 105 | - /// | ||
| 105 | + /// Resets `op_to_reset` with `op_or_function_name` and `raw_device_name`. This | ||
| 106 | + /// is for performance optimization by reusing an exiting unused op rather than | ||
| 107 | + /// creating a new op every time. If `raw_device_name` is `NULL` or empty, it | ||
| 108 | + /// does not set the device name. If it's not `NULL`, then it attempts to parse | ||
| 109 | + /// and set the device name. It's effectively `TFE_OpSetDevice`, but it is faster | ||
| 110 | + /// than separately calling it because if the existing op has the same | ||
| 111 | + /// `raw_device_name`, it skips parsing and just leave as it is. | ||
| 106 | 112 | /// </summary> | |
| 107 | - /// <param name="ctx">TFE_Context*</param> | ||
| 113 | + /// <param name="op_to_reset">TFE_Op*</param> | ||
| 108 | 114 | /// <param name="op_or_function_name">const char*</param> | |
| 115 | + /// <param name="raw_device_name">const char*</param> | ||
| 109 | 116 | /// <param name="status">TF_Status*</param> | |
| 110 | - /// <param name="op_to_reset">TFE_Op*</param> | ||
| 111 | 117 | [DllImport(TensorFlowLibName)] | |
| 112 | - public static extern void TFE_OpReset(IntPtr ctx, string op_or_function_name, IntPtr status, IntPtr op_to_reset); | ||
| 118 | + public static extern void TFE_OpReset(IntPtr op_to_reset, string op_or_function_name, string raw_device_name, IntPtr status); | ||
| 113 | 119 | ||
| 114 | 120 | /// <summary> | |
| 115 | 121 | /// | |
@@ -304,5 +310,17 @@ public partial class c_api | |||
| 304 | 310 | /// <returns>TFE_Executor*</returns> | |
| 305 | 311 | [DllImport(TensorFlowLibName)] | |
| 306 | 312 | public static extern TFE_Executor TFE_ContextGetExecutorForThread(IntPtr ctx); | |
| 313 | + | ||
| 314 | + [DllImport(TensorFlowLibName)] | ||
| 315 | + public static extern void TFE_Test(); | ||
| 316 | + | ||
| 317 | + [DllImport(TensorFlowLibName)] | ||
| 318 | + public static extern IntPtr TFE_TapeSetNew(bool persistent, bool watch_accessed_variables); | ||
| 319 | + | ||
| 320 | + [DllImport(TensorFlowLibName)] | ||
| 321 | + public static extern void TFE_TapeWatch(IntPtr tape, IntPtr tensor, int tensor_id); | ||
| 322 | + | ||
| 323 | + [DllImport(TensorFlowLibName)] | ||
| 324 | + public static extern void TFE_TapeGradient(IntPtr tape, long[] targetTensorIds, IntPtr[] target, long[] sourcesTensorIds, IntPtr status); | ||
| 307 | 325 | } | |
| 308 | 326 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,15 @@ | |||
| 1 | + using System.Collections.Generic; | ||
| 2 | + using System.Linq; | ||
| 3 | + using System; | ||
| 4 | + using static Tensorflow.OpDef.Types; | ||
| 5 | + | ||
| 6 | + namespace Tensorflow.Eager | ||
| 7 | + { | ||
| 8 | + /// <summary> | ||
| 9 | + /// python\eager\pywrap_tfe_src.cc | ||
| 10 | + /// </summary> | ||
| 11 | + public partial class wrap_tfe_src | ||
| 12 | + { | ||
| 13 | + | ||
| 14 | + } | ||
| 15 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -110,7 +110,7 @@ private static TFE_Op GetOp(Context ctx, string op_or_function_name, Status stat | |||
| 110 | 110 | var maybe_op = ReleaseThreadLocalOp(); | |
| 111 | 111 | if (maybe_op != IntPtr.Zero) | |
| 112 | 112 | { | |
| 113 | - c_api.TFE_OpReset(ctx, op_or_function_name, status, maybe_op); | ||
| 113 | + c_api.TFE_OpReset(maybe_op, op_or_function_name, ctx.device_name, status); | ||
| 114 | 114 | } | |
| 115 | 115 | else | |
| 116 | 116 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -23,7 +23,6 @@ public class GradientActor : IDisposable | |||
| 23 | 23 | bool _watch_accessed_variables; | |
| 24 | 24 | bool _created_eagerly; | |
| 25 | 25 | Tape _tape; | |
| 26 | - int tape_nesting_id_counter = 0; | ||
| 27 | 26 | ||
| 28 | 27 | public GradientActor(bool persistent = false, | |
| 29 | 28 | bool watch_accessed_variables = true) | |
@@ -41,18 +40,28 @@ private void _push_tape() | |||
| 41 | 40 | "re-enter an already-active tape."); | |
| 42 | 41 | ||
| 43 | 42 | if (_tape == null) | |
| 44 | - { | ||
| 45 | - _tape = new Tape(); | ||
| 46 | - _tape.tape = new GradientTape(_persistent, _watch_accessed_variables); | ||
| 47 | - _tape.nesting_id = tape_nesting_id_counter++; | ||
| 48 | - } | ||
| 43 | + _tape = new Tape(_persistent, _watch_accessed_variables); | ||
| 44 | + else | ||
| 45 | + throw new NotImplementedException(""); | ||
| 49 | 46 | ||
| 50 | 47 | _recording = true; | |
| 51 | 48 | } | |
| 52 | 49 | ||
| 50 | + /// <summary> | ||
| 51 | + /// Marks this tensor to be watched by the given tape. | ||
| 52 | + /// </summary> | ||
| 53 | + /// <param name="x"></param> | ||
| 53 | 54 | public void watch(Tensor x) | |
| 54 | 55 | { | |
| 56 | + _tape.watch(x); | ||
| 57 | + } | ||
| 55 | 58 | ||
| 59 | + public Tensor gradient(Tensor target, Tensor sources) | ||
| 60 | + { | ||
| 61 | + c_api.TFE_Test(); | ||
| 62 | + //using (var status = new Status()) | ||
| 63 | + //c_api.TFE_TapeGradient(_tape, new long[] { target.Id }, status); | ||
| 64 | + return null; | ||
| 56 | 65 | } | |
| 57 | 66 | ||
| 58 | 67 | public void Dispose() | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,11 +4,21 @@ | |||
| 4 | 4 | ||
| 5 | 5 | namespace Tensorflow.Gradients | |
| 6 | 6 | { | |
| 7 | - public class Tape | ||
| 7 | + public class Tape : DisposableObject | ||
| 8 | 8 | { | |
| 9 | 9 | public GradientTape tape { get; set; } | |
| 10 | 10 | public int nesting_id { get; set; } | |
| 11 | 11 | ||
| 12 | + public Tape(bool persistent, bool watch_accessed_variables) | ||
| 13 | + { | ||
| 14 | + _handle = c_api.TFE_TapeSetNew(persistent, watch_accessed_variables); | ||
| 15 | + } | ||
| 16 | + | ||
| 17 | + public void watch(Tensor x) | ||
| 18 | + { | ||
| 19 | + c_api.TFE_TapeWatch(_handle, x, x.Id); | ||
| 20 | + } | ||
| 21 | + | ||
| 12 | 22 | public static bool IsDtypeTrainable(DataType dtype) | |
| 13 | 23 | { | |
| 14 | 24 | switch (dtype) | |
@@ -26,5 +36,12 @@ public static bool IsDtypeTrainable(DataType dtype) | |||
| 26 | 36 | return false; | |
| 27 | 37 | } | |
| 28 | 38 | } | |
| 39 | + | ||
| 40 | + protected override void DisposeUnmanagedResources(IntPtr handle) | ||
| 41 | + { | ||
| 42 | + } | ||
| 43 | + | ||
| 44 | + public static implicit operator IntPtr(Tape tape) | ||
| 45 | + => tape._handle; | ||
| 29 | 46 | } | |
| 30 | 47 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,7 +39,7 @@ public partial class Tensor : DisposableObject, | |||
| 39 | 39 | IPackable<Tensor>, | |
| 40 | 40 | ICanBeFlattened | |
| 41 | 41 | { | |
| 42 | - private readonly int _id; | ||
| 42 | + protected int _id; | ||
| 43 | 43 | private readonly Operation _op; | |
| 44 | 44 | private readonly int _value_index; | |
| 45 | 45 | private TF_Output? _tf_output; | |
| Back | FazBrowse Home | New Git URL |
0 commit comments