| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 899d15e commit 13e4e3e
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -191,14 +191,13 @@ private static Tensor _constant_if_small(int value, Tensor shape) | |||
| 191 | 191 | ||
| 192 | 192 | private static Tensor _constant_if_small<T>(T value, TensorShape shape, TF_DataType dtype, string name) | |
| 193 | 193 | { | |
| 194 | - Tensor shape_t = null; | ||
| 195 | 194 | if (shape.size < 1000) | |
| 196 | 195 | { | |
| 197 | 196 | return constant_op.constant(value, shape: shape, dtype: dtype, name: name); | |
| 198 | 197 | } | |
| 199 | 198 | else | |
| 200 | 199 | { | |
| 201 | - shape_t = constant_op._tensor_shape_tensor_conversion_function(shape); | ||
| 200 | + var shape_t = constant_op._tensor_shape_tensor_conversion_function(shape); | ||
| 202 | 201 | var c = constant_op.constant(0, dtype: dtype); | |
| 203 | 202 | return gen_array_ops.fill(shape_t, c, name: name); | |
| 204 | 203 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -80,10 +80,11 @@ public RaggedTensor string_split_v2(Tensor input, string sep = " ", int maxsplit | |||
| 80 | 80 | var sep_tensor = ops.convert_to_tensor(sep, dtype: TF_DataType.TF_STRING); | |
| 81 | 81 | if(input.rank == 0) | |
| 82 | 82 | { | |
| 83 | - return string_split_v2(array_ops.stack(new[] { input }), | ||
| 83 | + var parts = string_split_v2(array_ops.stack(new[] { input }), | ||
| 84 | 84 | sep: sep, | |
| 85 | 85 | maxsplit: maxsplit, | |
| 86 | - name: name)[0]; | ||
| 86 | + name: name); | ||
| 87 | + return parts; | ||
| 87 | 88 | } | |
| 88 | 89 | ||
| 89 | 90 | var result = tf.Context.ExecuteOp("StringSplitV2", name, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -44,6 +44,18 @@ public TensorShape shape | |||
| 44 | 44 | } | |
| 45 | 45 | } | |
| 46 | 46 | ||
| 47 | + public Tensor this[int index] | ||
| 48 | + { | ||
| 49 | + get | ||
| 50 | + { | ||
| 51 | + return tf_with(ops.name_scope(null, "RaggedGetItem"), scope => | ||
| 52 | + { | ||
| 53 | + string name = scope; | ||
| 54 | + return _ragged_getitem(index); | ||
| 55 | + }); | ||
| 56 | + } | ||
| 57 | + } | ||
| 58 | + | ||
| 47 | 59 | public RaggedTensor this[params Slice[] slices] | |
| 48 | 60 | { | |
| 49 | 61 | get | |
@@ -61,6 +73,14 @@ public RaggedTensor this[params Slice[] slices] | |||
| 61 | 73 | } | |
| 62 | 74 | } | |
| 63 | 75 | ||
| 76 | + Tensor _ragged_getitem(int row_key) | ||
| 77 | + { | ||
| 78 | + var starts = _row_splits[":-1"]; | ||
| 79 | + var limits = _row_splits["1:"]; | ||
| 80 | + var row = _values[starts[row_key], limits[row_key]]; | ||
| 81 | + return row; | ||
| 82 | + } | ||
| 83 | + | ||
| 64 | 84 | RaggedTensor _ragged_getitem_inner_dimensions(RaggedTensor input, Slice[] slices) | |
| 65 | 85 | { | |
| 66 | 86 | return input; | |
@@ -134,7 +154,7 @@ Tensor[] nested_row_splits | |||
| 134 | 154 | => new[] { _row_splits }; | |
| 135 | 155 | ||
| 136 | 156 | public override string ToString() | |
| 137 | - => $"tf.RaggedTensor: shape={_values.TensorShape} [{string.Join(", ", _values.StringData().Take(10))}]"; | ||
| 157 | + => $"tf.RaggedTensor: shape={shape} [{string.Join(", ", _values.StringData().Take(10))}]"; | ||
| 138 | 158 | ||
| 139 | 159 | public static implicit operator Tensor(RaggedTensor indexedSlices) | |
| 140 | 160 | => indexedSlices._to_variant(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -120,11 +120,11 @@ public Tensor slice(Slice slice) | |||
| 120 | 120 | }); | |
| 121 | 121 | } | |
| 122 | 122 | ||
| 123 | - public Tensor this[params Tensor[] slice] | ||
| 123 | + public Tensor this[Tensor start, Tensor stop = null, Tensor step = null] | ||
| 124 | 124 | { | |
| 125 | 125 | get | |
| 126 | 126 | { | |
| 127 | - var args = tensor_util.ParseSlices(slice); | ||
| 127 | + var args = tensor_util.ParseSlices(start, stop: stop, step: step); | ||
| 128 | 128 | ||
| 129 | 129 | return tf_with(ops.name_scope(null, "strided_slice", args), scope => | |
| 130 | 130 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -674,25 +674,30 @@ public static ParsedSliceArgs ParseSlices(Slice[] slices) | |||
| 674 | 674 | }; | |
| 675 | 675 | } | |
| 676 | 676 | ||
| 677 | - public static ParsedSliceArgs ParseSlices(Tensor[] slices) | ||
| 677 | + public static ParsedSliceArgs ParseSlices(Tensor start, Tensor stop = null, Tensor step = null) | ||
| 678 | 678 | { | |
| 679 | 679 | var begin = new List<Tensor>(); | |
| 680 | 680 | var end = new List<Tensor>(); | |
| 681 | 681 | var strides = new List<Tensor>(); | |
| 682 | 682 | ||
| 683 | - var index = 0; | ||
| 683 | + // var index = 0; | ||
| 684 | 684 | var (new_axis_mask, shrink_axis_mask) = (0, 0); | |
| 685 | 685 | var (begin_mask, end_mask) = (0, 0); | |
| 686 | 686 | var ellipsis_mask = 0; | |
| 687 | 687 | ||
| 688 | - foreach (var s in slices) | ||
| 689 | - { | ||
| 690 | - begin.Add(s); | ||
| 691 | - end.Add(s + 1); | ||
| 692 | - shrink_axis_mask |= (1 << index); | ||
| 693 | - strides.Add(tf.constant(1, dtype: s.dtype)); | ||
| 694 | - index += 1; | ||
| 695 | - } | ||
| 688 | + begin.Add(start); | ||
| 689 | + | ||
| 690 | + if (stop == null) | ||
| 691 | + end.Add(start + 1); | ||
| 692 | + else | ||
| 693 | + end.Add(stop); | ||
| 694 | + | ||
| 695 | + // shrink_axis_mask |= (1 << index); | ||
| 696 | + | ||
| 697 | + if (step == null) | ||
| 698 | + strides.Add(tf.constant(1, dtype: start.dtype)); | ||
| 699 | + else | ||
| 700 | + strides.Add(step); | ||
| 696 | 701 | ||
| 697 | 702 | return new ParsedSliceArgs | |
| 698 | 703 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -40,6 +40,9 @@ public Tensor constant(object value, | |||
| 40 | 40 | public Tensor zeros(TensorShape shape, TF_DataType dtype = TF_DataType.TF_FLOAT, string name = null) | |
| 41 | 41 | => array_ops.zeros(shape, dtype, name); | |
| 42 | 42 | ||
| 43 | + public Tensor zeros(Tensor shape, TF_DataType dtype = TF_DataType.TF_FLOAT, string name = null) | ||
| 44 | + => array_ops.zeros(shape, dtype, name); | ||
| 45 | + | ||
| 43 | 46 | public Tensor ones(TensorShape shape, TF_DataType dtype = TF_DataType.TF_FLOAT, string name = null) | |
| 44 | 47 | => array_ops.ones(shape, dtype, name); | |
| 45 | 48 | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments