| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 71fef9a commit afaf0c8
1 file changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,6 +10,7 @@ namespace Tensorflow.Eager | |||
| 10 | 10 | /// </summary> | |
| 11 | 11 | public class pywrap_tfe_src | |
| 12 | 12 | { | |
| 13 | + static int kFastPathExecuteInputStartIndex = 0; | ||
| 13 | 14 | public static EagerTensor TFE_Py_FastPathExecute(Context ctx, | |
| 14 | 15 | string device_name, | |
| 15 | 16 | string opName, | |
@@ -28,7 +29,7 @@ public static EagerTensor TFE_Py_FastPathExecute(Context ctx, | |||
| 28 | 29 | ||
| 29 | 30 | // Set non-inferred attrs, including setting defaults if the attr is passed in | |
| 30 | 31 | // as None. | |
| 31 | - for (int i = op_def.InputArg.Count; i < args_size; i += 2) | ||
| 32 | + for (int i = kFastPathExecuteInputStartIndex + op_def.InputArg.Count; i < args_size; i += 2) | ||
| 32 | 33 | { | |
| 33 | 34 | var attr_name = args[i].ToString(); | |
| 34 | 35 | var attr_value = args[i + 1]; | |
@@ -38,20 +39,39 @@ public static EagerTensor TFE_Py_FastPathExecute(Context ctx, | |||
| 38 | 39 | if(attr_name == attr.Name) | |
| 39 | 40 | { | |
| 40 | 41 | SetOpAttrWithDefaults(ctx, op, attr, attr_name, attr_value, attr_list_sizes, status); | |
| 42 | + status.Check(true); | ||
| 41 | 43 | break; | |
| 42 | 44 | } | |
| 43 | 45 | } | |
| 44 | 46 | } | |
| 45 | 47 | ||
| 46 | 48 | c_api.TFE_OpSetDevice(op, device_name, status); | |
| 49 | + status.Check(true); | ||
| 47 | 50 | ||
| 51 | + // Add inferred attrs and inputs. | ||
| 48 | 52 | for (int i = 0; i < op_def.InputArg.Count; i++) | |
| 49 | 53 | { | |
| 50 | 54 | var input_arg = op_def.InputArg[i]; | |
| 55 | + int len = (args[kFastPathExecuteInputStartIndex + i] as object[]).Length; | ||
| 51 | 56 | if (!string.IsNullOrEmpty(input_arg.NumberAttr)) | |
| 52 | 57 | { | |
| 53 | - c_api.TFE_OpSetAttrInt(op, input_arg.NumberAttr, 0); | ||
| 54 | - attr_list_sizes[input_arg.NumberAttr] = 0; | ||
| 58 | + c_api.TFE_OpSetAttrInt(op, input_arg.NumberAttr, len); | ||
| 59 | + attr_list_sizes[input_arg.NumberAttr] = len; | ||
| 60 | + | ||
| 61 | + if (len > 0) | ||
| 62 | + { | ||
| 63 | + var fast_input_array = (object[])args[i]; | ||
| 64 | + // First item adds the type attr. | ||
| 65 | + if (!AddInputToOp(fast_input_array[i], true, input_arg, op, status)) | ||
| 66 | + return null; | ||
| 67 | + | ||
| 68 | + for (var j = 1; j < len; j++) | ||
| 69 | + { | ||
| 70 | + // Since the list is homogeneous, we don't need to re-add the attr. | ||
| 71 | + if (!AddInputToOp(fast_input_array[j], false, input_arg, op, status)) | ||
| 72 | + return null; | ||
| 73 | + } | ||
| 74 | + } | ||
| 55 | 75 | } | |
| 56 | 76 | else if (!string.IsNullOrEmpty(input_arg.TypeListAttr)) | |
| 57 | 77 | { | |
@@ -60,14 +80,7 @@ public static EagerTensor TFE_Py_FastPathExecute(Context ctx, | |||
| 60 | 80 | else | |
| 61 | 81 | { | |
| 62 | 82 | // The item is a single item. | |
| 63 | - switch (args[i]) | ||
| 64 | - { | ||
| 65 | - case Tensor inputTensor: | ||
| 66 | - AddInputToOp(inputTensor, true, input_arg, op, status); | ||
| 67 | - break; | ||
| 68 | - default: | ||
| 69 | - throw new NotImplementedException(""); | ||
| 70 | - } | ||
| 83 | + AddInputToOp(args[i], true, input_arg, op, status); | ||
| 71 | 84 | } | |
| 72 | 85 | } | |
| 73 | 86 | ||
@@ -106,13 +119,23 @@ public static EagerTensor TFE_Py_FastPathExecute(Context ctx, | |||
| 106 | 119 | /// <param name="op"></param> | |
| 107 | 120 | /// <param name="status"></param> | |
| 108 | 121 | /// <returns></returns> | |
| 109 | - private static bool AddInputToOp(Tensor input, | ||
| 122 | + private static bool AddInputToOp(object inputs, | ||
| 110 | 123 | bool add_type_attr, | |
| 111 | 124 | ArgDef input_arg, | |
| 112 | 125 | IntPtr op, | |
| 113 | 126 | Status status) | |
| 114 | 127 | { | |
| 115 | - var input_handle = c_api.TFE_NewTensorHandle(input, status); | ||
| 128 | + IntPtr input_handle = IntPtr.Zero; | ||
| 129 | + | ||
| 130 | + switch (inputs) | ||
| 131 | + { | ||
| 132 | + case Tensor input: | ||
| 133 | + input_handle = c_api.TFE_NewTensorHandle(input, status); | ||
| 134 | + break; | ||
| 135 | + default: | ||
| 136 | + throw new NotImplementedException(""); | ||
| 137 | + } | ||
| 138 | + | ||
| 116 | 139 | ||
| 117 | 140 | if(add_type_attr && !string.IsNullOrEmpty(input_arg.TypeAttr)) | |
| 118 | 141 | { | |
| Back | FazBrowse Home | New Git URL |
0 commit comments