| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 3f8b658 commit ca6f8b2
15 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -390,7 +390,7 @@ public Tensor divide<T>(Tensor x, T[] y, string name = null) where T : struct | |||
| 390 | 390 | => x / ops.convert_to_tensor(y, dtype: x.dtype.as_base_dtype(), name: "y"); | |
| 391 | 391 | ||
| 392 | 392 | public Tensor pow<T1, T2>(T1 x, T2 y, string name = "pow") | |
| 393 | - => gen_math_ops.pow(x, y, name: name); | ||
| 393 | + => math_ops.pow(x, y, name: name); | ||
| 394 | 394 | ||
| 395 | 395 | /// <summary> | |
| 396 | 396 | /// Divides `x / y` elementwise, rounding toward the most negative integer. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -53,6 +53,9 @@ public override string ToString() | |||
| 53 | 53 | ||
| 54 | 54 | public static string GetFormattedString(TF_DataType dtype, NDArray nd) | |
| 55 | 55 | { | |
| 56 | + if (nd.size == 0) | ||
| 57 | + return "[]"; | ||
| 58 | + | ||
| 56 | 59 | switch (dtype) | |
| 57 | 60 | { | |
| 58 | 61 | case TF_DataType.TF_STRING: | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -375,6 +375,9 @@ public static extern IntPtr TFE_QuickExecute(IntPtr ctx, | |||
| 375 | 375 | [DllImport(TensorFlowLibName)] | |
| 376 | 376 | public static extern void TFE_TapeWatch(IntPtr tape, IntPtr tensor); | |
| 377 | 377 | ||
| 378 | + [DllImport(TensorFlowLibName)] | ||
| 379 | + public static extern void TFE_TapeVariableAccessed(IntPtr variable); | ||
| 380 | + | ||
| 378 | 381 | [DllImport(TensorFlowLibName)] | |
| 379 | 382 | public static extern IntPtr TFE_TapeGradient(IntPtr tape, | |
| 380 | 383 | IntPtr[] target, int target_size, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,6 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | + using System.Linq; | ||
| 3 | 4 | using System.Text; | |
| 4 | 5 | using Tensorflow.Eager; | |
| 5 | 6 | using static Tensorflow.Binding; | |
@@ -65,7 +66,7 @@ public void watch(Tensor x) | |||
| 65 | 66 | _tape.watch(x as EagerTensor); | |
| 66 | 67 | } | |
| 67 | 68 | ||
| 68 | - public Tensor gradient(Tensor target, Tensor sources) | ||
| 69 | + public Tensor gradient(Tensor target, Tensor source) | ||
| 69 | 70 | { | |
| 70 | 71 | if(_recording) | |
| 71 | 72 | { | |
@@ -76,15 +77,33 @@ public Tensor gradient(Tensor target, Tensor sources) | |||
| 76 | 77 | using var status = new Status(); | |
| 77 | 78 | var et = c_api.TFE_TapeGradient(_tape, | |
| 78 | 79 | new [] { (target as EagerTensor).EagerTensorHandle }, 1, | |
| 79 | - new [] { (sources as EagerTensor).EagerTensorHandle }, 1, | ||
| 80 | + new [] { (source as EagerTensor).EagerTensorHandle }, 1, | ||
| 80 | 81 | status); | |
| 81 | 82 | status.Check(true); | |
| 82 | 83 | return new EagerTensor(et); | |
| 83 | 84 | } | |
| 84 | 85 | ||
| 86 | + public Tensor gradient(Tensor target, ResourceVariable[] sources) | ||
| 87 | + { | ||
| 88 | + if (_recording) | ||
| 89 | + { | ||
| 90 | + if (!_persistent) | ||
| 91 | + _pop_tape(); | ||
| 92 | + } | ||
| 93 | + | ||
| 94 | + using var status = new Status(); | ||
| 95 | + EagerTensorHandle et = c_api.TFE_TapeGradient(_tape, | ||
| 96 | + new[] { (target as EagerTensor).EagerTensorHandle }, 1, | ||
| 97 | + sources.Select(x => (x.handle as EagerTensor).EagerTensorHandle).ToArray(), sources.Length, | ||
| 98 | + status); | ||
| 99 | + status.Check(true); | ||
| 100 | + return et; | ||
| 101 | + } | ||
| 102 | + | ||
| 85 | 103 | public void Dispose() | |
| 86 | 104 | { | |
| 87 | - | ||
| 105 | + if (_recording) | ||
| 106 | + _pop_tape(); | ||
| 88 | 107 | } | |
| 89 | 108 | } | |
| 90 | 109 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -25,6 +25,11 @@ public void pop_tape(Tape tape) | |||
| 25 | 25 | c_api.TFE_TapeSetRemove(tape); | |
| 26 | 26 | } | |
| 27 | 27 | ||
| 28 | + public static void variable_accessed(ResourceVariable variable) | ||
| 29 | + { | ||
| 30 | + c_api.TFE_TapeVariableAccessed(variable.handle as EagerTensor); | ||
| 31 | + } | ||
| 32 | + | ||
| 28 | 33 | public static bool IsDtypeTrainable(DataType dtype) | |
| 29 | 34 | { | |
| 30 | 35 | switch (dtype) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -220,6 +220,18 @@ public static Tensor prevent_gradient(Tensor input, string message = "", string | |||
| 220 | 220 | /// <param name="name"></param> | |
| 221 | 221 | public static Tensor identity(Tensor input, string name = null) | |
| 222 | 222 | { | |
| 223 | + if (tf.context.executing_eagerly()) | ||
| 224 | + { | ||
| 225 | + using var status = new Status(); | ||
| 226 | + EagerTensorHandle tensor = c_api.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 227 | + "Identity", name, new IntPtr[] | ||
| 228 | + { | ||
| 229 | + input as EagerTensor | ||
| 230 | + }, 1, null, status); | ||
| 231 | + status.Check(true); | ||
| 232 | + return tensor; | ||
| 233 | + } | ||
| 234 | + | ||
| 223 | 235 | var _op = _op_def_lib._apply_op_helper("Identity", name, new { input }); | |
| 224 | 236 | ||
| 225 | 237 | return _op.output; | |
@@ -258,14 +270,14 @@ public static Tensor fill<T>(Tensor dims, T value, string name = null) | |||
| 258 | 270 | if (tf.context.executing_eagerly()) | |
| 259 | 271 | { | |
| 260 | 272 | using var status = new Status(); | |
| 261 | - var tensor = c_api.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 273 | + EagerTensorHandle tensor = c_api.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 262 | 274 | "Fill", name, new IntPtr[] | |
| 263 | 275 | { | |
| 264 | 276 | dims as EagerTensor, | |
| 265 | 277 | value as EagerTensor | |
| 266 | 278 | }, 2, null, status); | |
| 267 | 279 | status.Check(true); | |
| 268 | - return new EagerTensor(tensor); | ||
| 280 | + return tensor; | ||
| 269 | 281 | } | |
| 270 | 282 | ||
| 271 | 283 | var _op = _op_def_lib._apply_op_helper("Fill", name, new { dims, value }); | |
@@ -281,6 +293,18 @@ value as EagerTensor | |||
| 281 | 293 | /// <returns>A tuple of `Tensor` objects (r0, r1).</returns> | |
| 282 | 294 | public static (Tensor, Tensor) broadcast_gradient_args(Tensor s0, Tensor s1, string name = "") | |
| 283 | 295 | { | |
| 296 | + if (tf.context.executing_eagerly()) | ||
| 297 | + { | ||
| 298 | + using var status = new Status(); | ||
| 299 | + var _result = c_api.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 300 | + "BroadcastGradientArgs", name, new IntPtr[] | ||
| 301 | + { | ||
| 302 | + s0 as EagerTensor, | ||
| 303 | + s1 as EagerTensor | ||
| 304 | + }, 2, null, status); | ||
| 305 | + status.Check(true); | ||
| 306 | + } | ||
| 307 | + | ||
| 284 | 308 | var _op = _op_def_lib._apply_op_helper("BroadcastGradientArgs", name, new { s0, s1 }); | |
| 285 | 309 | ||
| 286 | 310 | return (_op.outputs[0], _op.outputs[1]); | |
@@ -371,10 +395,19 @@ public static Tensor shape(Tensor input, TF_DataType out_type = TF_DataType.TF_I | |||
| 371 | 395 | { | |
| 372 | 396 | if (tf.context.executing_eagerly()) | |
| 373 | 397 | { | |
| 374 | - var _result = wrap_tfe_src.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 375 | - "Shape", name, null, | ||
| 376 | - input, "out_type", out_type); | ||
| 377 | - return _result; | ||
| 398 | + using var status = new Status(); | ||
| 399 | + EagerTensorHandle tensor = c_api.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 400 | + "Shape", name, new IntPtr[] | ||
| 401 | + { | ||
| 402 | + input as EagerTensor, | ||
| 403 | + }, 1, | ||
| 404 | + op => wrap_tfe_src.SetOpAttrs(tf.context, op, new object[] | ||
| 405 | + { | ||
| 406 | + "out_type", out_type | ||
| 407 | + }, status), | ||
| 408 | + status); | ||
| 409 | + status.Check(true); | ||
| 410 | + return tensor; | ||
| 378 | 411 | } | |
| 379 | 412 | ||
| 380 | 413 | var _op = _op_def_lib._apply_op_helper("Shape", name, new { input, out_type }); | |
@@ -455,12 +488,26 @@ public static Tensor strided_slice(Tensor input, Tensor begin, Tensor end, Tenso | |||
| 455 | 488 | { | |
| 456 | 489 | if (tf.context.executing_eagerly()) | |
| 457 | 490 | { | |
| 458 | - var _result = wrap_tfe_src.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 459 | - "StridedSlice", name, null, | ||
| 460 | - input, begin, end, strides, "begin_mask", begin_mask, | ||
| 461 | - "end_mask", end_mask, "ellipsis_mask", ellipsis_mask, | ||
| 462 | - "new_axis_mask", new_axis_mask, "shrink_axis_mask", shrink_axis_mask); | ||
| 463 | - return _result; | ||
| 491 | + using var status = new Status(); | ||
| 492 | + EagerTensorHandle tensor = c_api.TFE_FastPathExecute(tf.context, tf.context.device_name, | ||
| 493 | + "StridedSlice", name, new IntPtr[] | ||
| 494 | + { | ||
| 495 | + input as EagerTensor, | ||
| 496 | + begin as EagerTensor, | ||
| 497 | + end as EagerTensor, | ||
| 498 | + strides as EagerTensor, | ||
| 499 | + }, 4, | ||
| 500 | + op => wrap_tfe_src.SetOpAttrs(tf.context, op, new object[] | ||
| 501 | + { | ||
| 502 | + "begin_mask", begin_mask, | ||
| 503 | + "end_mask", end_mask, | ||
| 504 | + "ellipsis_mask", ellipsis_mask, | ||
| 505 | + "new_axis_mask", new_axis_mask, | ||
| 506 | + "shrink_axis_mask", shrink_axis_mask | ||
| 507 | + }, status), | ||
| 508 | + status); | ||
| 509 | + status.Check(true); | ||
| 510 | + return tensor; | ||
| 464 | 511 | } | |
| 465 | 512 | ||
| 466 | 513 | var _op = _op_def_lib._apply_op_helper("StridedSlice", name, new | |
| Back | FazBrowse Home | New Git URL |
0 commit comments