| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 70f873e commit ed1a8d2
3 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -603,7 +603,17 @@ public static Tensor shape_internal(Tensor input, string name = null, bool optim | |||
| 603 | 603 | } | |
| 604 | 604 | } | |
| 605 | 605 | ||
| 606 | - return gen_array_ops.shape(input, name: name, out_type: out_type); | ||
| 606 | + return tf.Context.ExecuteOp("Shape", name, new ExecuteOpArgs(input) | ||
| 607 | + { | ||
| 608 | + GetGradientAttrs = (op) => new | ||
| 609 | + { | ||
| 610 | + T = op.get_attr<TF_DataType>("T"), | ||
| 611 | + out_type = op.get_attr<TF_DataType>("out_type") | ||
| 612 | + } | ||
| 613 | + }.SetAttributes(new | ||
| 614 | + { | ||
| 615 | + out_type | ||
| 616 | + })).First(); | ||
| 607 | 617 | }); | |
| 608 | 618 | } | |
| 609 | 619 | ||
@@ -703,23 +713,26 @@ public static Tensor strided_slice(Tensor input_, Tensor begin, Tensor end, | |||
| 703 | 713 | int new_axis_mask = 0, | |
| 704 | 714 | int shrink_axis_mask = 0, | |
| 705 | 715 | string name = null) | |
| 706 | - { | ||
| 707 | - var op = gen_array_ops.strided_slice( | ||
| 708 | - input: input_, | ||
| 709 | - begin: begin, | ||
| 710 | - end: end, | ||
| 711 | - strides: strides, | ||
| 712 | - begin_mask: begin_mask, | ||
| 713 | - end_mask: end_mask, | ||
| 714 | - ellipsis_mask: ellipsis_mask, | ||
| 715 | - new_axis_mask: new_axis_mask, | ||
| 716 | - shrink_axis_mask: shrink_axis_mask, | ||
| 717 | - name: name); | ||
| 718 | - | ||
| 719 | - string parent_name = name; | ||
| 720 | - | ||
| 721 | - return op; | ||
| 722 | - } | ||
| 716 | + => tf.Context.ExecuteOp("StridedSlice", name, new ExecuteOpArgs(input_, begin, end, strides) | ||
| 717 | + { | ||
| 718 | + GetGradientAttrs = (op) => new | ||
| 719 | + { | ||
| 720 | + T = op.get_attr<TF_DataType>("T"), | ||
| 721 | + Index = op.get_attr<TF_DataType>("Index"), | ||
| 722 | + begin_mask = op.get_attr<long>("begin_mask"), | ||
| 723 | + end_mask = op.get_attr<long>("end_mask"), | ||
| 724 | + ellipsis_mask = op.get_attr<long>("ellipsis_mask"), | ||
| 725 | + new_axis_mask = op.get_attr<long>("new_axis_mask"), | ||
| 726 | + shrink_axis_mask = op.get_attr<long>("shrink_axis_mask") | ||
| 727 | + } | ||
| 728 | + }.SetAttributes(new | ||
| 729 | + { | ||
| 730 | + begin_mask, | ||
| 731 | + end_mask, | ||
| 732 | + ellipsis_mask, | ||
| 733 | + new_axis_mask, | ||
| 734 | + shrink_axis_mask | ||
| 735 | + })); | ||
| 723 | 736 | ||
| 724 | 737 | /// <summary> | |
| 725 | 738 | /// Returns the gradient of `StridedSlice`. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,12 +5,17 @@ | |||
| 5 | 5 | /// </summary> | |
| 6 | 6 | public class KerasTensor | |
| 7 | 7 | { | |
| 8 | - private Tensor _tensor; | ||
| 9 | - public void SetTensor(Tensors tensor) | ||
| 10 | - => _tensor = tensor; | ||
| 8 | + private Tensors _inferred_value; | ||
| 9 | + public Tensors inferred_value | ||
| 10 | + { | ||
| 11 | + get => _inferred_value; | ||
| 12 | + set => _inferred_value = value; | ||
| 13 | + } | ||
| 11 | 14 | ||
| 12 | - private TensorSpec _type_spec; | ||
| 13 | 15 | private string _name; | |
| 16 | + private TensorSpec _type_spec; | ||
| 17 | + public Shape shape => _type_spec.shape; | ||
| 18 | + public TF_DataType dtype => _type_spec.dtype; | ||
| 14 | 19 | ||
| 15 | 20 | public KerasTensor(TensorSpec type_spec, string name = null) | |
| 16 | 21 | { | |
@@ -22,15 +27,23 @@ public static KerasTensor from_tensor(Tensor tensor) | |||
| 22 | 27 | { | |
| 23 | 28 | var type_spec = tensor.ToTensorSpec(); | |
| 24 | 29 | var kt = new KerasTensor(type_spec, name: tensor.name); | |
| 25 | - kt.SetTensor(tensor); | ||
| 30 | + kt.inferred_value = tensor; | ||
| 26 | 31 | return kt; | |
| 27 | 32 | } | |
| 28 | 33 | ||
| 34 | + public override string ToString() | ||
| 35 | + => _inferred_value.Length switch | ||
| 36 | + { | ||
| 37 | + > 1 => "[" + string.Join(", ", _inferred_value.Select(x => $"<KerasTensor: shape={x.shape} dtype={x.dtype}>")) + "]", | ||
| 38 | + 1 => $"<KerasTensor: shape={_inferred_value.shape} dtype={_inferred_value.dtype}>", | ||
| 39 | + _ => _inferred_value.ToString(), | ||
| 40 | + }; | ||
| 41 | + | ||
| 29 | 42 | public static implicit operator Tensors(KerasTensor kt) | |
| 30 | - => kt._tensor; | ||
| 43 | + => kt._inferred_value; | ||
| 31 | 44 | ||
| 32 | 45 | public static implicit operator Tensor(KerasTensor kt) | |
| 33 | - => kt._tensor; | ||
| 46 | + => kt._inferred_value; | ||
| 34 | 47 | ||
| 35 | 48 | public static implicit operator KerasTensor(Tensor tensor) | |
| 36 | 49 | => from_tensor(tensor); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -42,7 +42,7 @@ public Tensor this[params Slice[] slices] | |||
| 42 | 42 | array_ops.stack(args.End), | |
| 43 | 43 | array_ops.stack(args.Strides)); | |
| 44 | 44 | ||
| 45 | - return gen_array_ops.strided_slice( | ||
| 45 | + return array_ops.strided_slice( | ||
| 46 | 46 | this, | |
| 47 | 47 | packed_begin, | |
| 48 | 48 | packed_end, | |
| Back | FazBrowse Home | New Git URL |
0 commit comments