| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent bd154e8 commit bdf229a
7 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -80,6 +80,14 @@ public static NDArray MakeNdarray(TensorProto tensor) | |||
| 80 | 80 | { | |
| 81 | 81 | return np.array(tensor.IntVal.ToArray()).reshape(shape); | |
| 82 | 82 | } | |
| 83 | + else if (new DataType[] { DataType.DtInt64 }.Contains(tensor.Dtype)) | ||
| 84 | + { | ||
| 85 | + return np.array(tensor.Int64Val.ToArray()).reshape(shape); | ||
| 86 | + } | ||
| 87 | + else if (new DataType[] { DataType.DtUint64 }.Contains(tensor.Dtype)) | ||
| 88 | + { | ||
| 89 | + return np.array(tensor.Uint64Val.ToArray()).reshape(shape); | ||
| 90 | + } | ||
| 83 | 91 | else if (tensor.Dtype == DataType.DtBool) | |
| 84 | 92 | { | |
| 85 | 93 | return np.array(tensor.BoolVal.ToArray()).reshape(shape); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,18 @@ | |||
| 1 | + using Google.Protobuf.Collections; | ||
| 2 | + using System.IO; | ||
| 3 | + using Tensorflow.Train; | ||
| 4 | + | ||
| 5 | + namespace Tensorflow.Trackables; | ||
| 6 | + | ||
| 7 | + public class AssetResource : Trackable | ||
| 8 | + { | ||
| 9 | + public static (Trackable, Action<object, object, object>) deserialize_from_proto(SavedObject object_proto, | ||
| 10 | + string export_dir, | ||
| 11 | + RepeatedField<AssetFileDef> asset_file_def, | ||
| 12 | + Dictionary<string, MapField<string, AttrValue>> operation_attributes) | ||
| 13 | + { | ||
| 14 | + var proto = object_proto.Asset; | ||
| 15 | + var filename = Path.Combine(export_dir, asset_file_def[proto.AssetFileDefIndex].Filename); | ||
| 16 | + return (new AssetResource(), null); | ||
| 17 | + } | ||
| 18 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,12 +1,13 @@ | |||
| 1 | - using System.Runtime.CompilerServices; | ||
| 1 | + using Google.Protobuf.Collections; | ||
| 2 | 2 | using Tensorflow.Train; | |
| 3 | 3 | ||
| 4 | 4 | namespace Tensorflow.Trackables; | |
| 5 | 5 | ||
| 6 | 6 | public class RestoredResource : TrackableResource | |
| 7 | 7 | { | |
| 8 | - public static (Trackable, Action<object, object, object>) deserialize_from_proto() | ||
| 8 | + public static (Trackable, Action<object, object, object>) deserialize_from_proto(SavedObject object_proto, | ||
| 9 | + Dictionary<string, MapField<string, AttrValue>> operation_attributes) | ||
| 9 | 10 | { | |
| 10 | - return (null, null); | ||
| 11 | + return (new RestoredResource(), null); | ||
| 11 | 12 | } | |
| 12 | 13 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,11 +1,22 @@ | |||
| 1 | - using Tensorflow.Train; | ||
| 1 | + using Google.Protobuf.Collections; | ||
| 2 | + using Tensorflow.Train; | ||
| 2 | 3 | ||
| 3 | 4 | namespace Tensorflow.Trackables; | |
| 4 | 5 | ||
| 5 | 6 | public class TrackableConstant : Trackable | |
| 6 | 7 | { | |
| 7 | - public static (Trackable, Action<object, object, object>) deserialize_from_proto() | ||
| 8 | + Tensor _constant; | ||
| 9 | + public TrackableConstant(Tensor constant) | ||
| 8 | 10 | { | |
| 9 | - return (null, null); | ||
| 11 | + _constant = constant; | ||
| 12 | + } | ||
| 13 | + | ||
| 14 | + public static (Trackable, Action<object, object, object>) deserialize_from_proto(SavedObject object_proto, | ||
| 15 | + Dictionary<string, MapField<string, AttrValue>> operation_attributes) | ||
| 16 | + { | ||
| 17 | + var tensor_proto = operation_attributes[object_proto.Constant.Operation]["value"].Tensor; | ||
| 18 | + var ndarray = tensor_util.MakeNdarray(tensor_proto); | ||
| 19 | + var imported_constant = constant_op.constant(ndarray); | ||
| 20 | + return (new TrackableConstant(imported_constant), null); | ||
| 10 | 21 | } | |
| 11 | 22 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -9,6 +9,19 @@ namespace Tensorflow.Training.Saving.SavedModel | |||
| 9 | 9 | { | |
| 10 | 10 | public static class function_deserialization | |
| 11 | 11 | { | |
| 12 | + /// <summary> | ||
| 13 | + /// Creates a `Function` from a `SavedFunction`. | ||
| 14 | + /// </summary> | ||
| 15 | + /// <param name="saved_concrete_function"></param> | ||
| 16 | + /// <param name="concrete_functions"></param> | ||
| 17 | + /// <returns></returns> | ||
| 18 | + public static ConcreteFunction recreate_function(SavedFunction saved_concrete_function, | ||
| 19 | + IDictionary<string, ConcreteFunction> concrete_functions) | ||
| 20 | + { | ||
| 21 | + var function_spec = _deserialize_function_spec_as_nonmethod(saved_concrete_function.FunctionSpec); | ||
| 22 | + return null; | ||
| 23 | + } | ||
| 24 | + | ||
| 12 | 25 | public static ConcreteFunction setup_bare_concrete_function(SavedBareConcreteFunction saved_bare_concrete_function, | |
| 13 | 26 | IDictionary<string, ConcreteFunction> concrete_functions) | |
| 14 | 27 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -387,13 +387,6 @@ private void _load_nodes() | |||
| 387 | 387 | } | |
| 388 | 388 | else | |
| 389 | 389 | { | |
| 390 | - // skip the function and concrete function. | ||
| 391 | - if(proto.KindCase == SavedObject.KindOneofCase.BareConcreteFunction || proto.KindCase == SavedObject.KindOneofCase.Function) | ||
| 392 | - { | ||
| 393 | - nodes[node_id] = null; | ||
| 394 | - node_setters[node_id] = null; | ||
| 395 | - continue; | ||
| 396 | - } | ||
| 397 | 390 | var (node, setter) = _recreate(proto, node_id, nodes); | |
| 398 | 391 | nodes[node_id] = node; | |
| 399 | 392 | node_setters[node_id] = setter; | |
@@ -471,6 +464,11 @@ private void _load_edges() | |||
| 471 | 464 | } | |
| 472 | 465 | } | |
| 473 | 466 | ||
| 467 | + private void _setup_function_captures() | ||
| 468 | + { | ||
| 469 | + // TODO: implement it with concrete functions. | ||
| 470 | + } | ||
| 471 | + | ||
| 474 | 472 | private void _setup_remaining_functions() | |
| 475 | 473 | { | |
| 476 | 474 | // TODO: implement it with concrete functions. | |
@@ -542,9 +540,9 @@ private void _add_object_graph_edges(SavedObject proto, int node_id) | |||
| 542 | 540 | ||
| 543 | 541 | return proto.KindCase switch | |
| 544 | 542 | { | |
| 545 | - SavedObject.KindOneofCase.Resource => RestoredResource.deserialize_from_proto(), | ||
| 546 | - SavedObject.KindOneofCase.Asset => Asset.deserialize_from_proto(), | ||
| 547 | - SavedObject.KindOneofCase.Constant => TrackableConstant.deserialize_from_proto(), | ||
| 543 | + SavedObject.KindOneofCase.Resource => RestoredResource.deserialize_from_proto(proto, _operation_attributes), | ||
| 544 | + SavedObject.KindOneofCase.Asset => AssetResource.deserialize_from_proto(proto, _export_dir, _asset_file_def, _operation_attributes), | ||
| 545 | + SavedObject.KindOneofCase.Constant => TrackableConstant.deserialize_from_proto(proto, _operation_attributes), | ||
| 548 | 546 | _ => _recreate_default(proto, node_id, dependencies) | |
| 549 | 547 | }; | |
| 550 | 548 | } | |
@@ -563,7 +561,8 @@ private void _add_object_graph_edges(SavedObject proto, int node_id) | |||
| 563 | 561 | SavedObject.KindOneofCase.Function => _recreate_function(proto.Function, null), | |
| 564 | 562 | SavedObject.KindOneofCase.BareConcreteFunction => throw new NotImplementedException(), | |
| 565 | 563 | SavedObject.KindOneofCase.Variable => _recreate_variable(proto.Variable), | |
| 566 | - SavedObject.KindOneofCase.CapturedTensor => throw new NotImplementedException() | ||
| 564 | + SavedObject.KindOneofCase.CapturedTensor => throw new NotImplementedException(), | ||
| 565 | + _ => throw new NotImplementedException() | ||
| 567 | 566 | }; | |
| 568 | 567 | } | |
| 569 | 568 | ||
@@ -623,8 +622,12 @@ private void _add_object_graph_edges(SavedObject proto, int node_id) | |||
| 623 | 622 | private (ConcreteFunction, Action<object, object, object>) _recreate_function(SavedFunction proto, | |
| 624 | 623 | Dictionary<Maybe<string, int>, Trackable> dependencies) | |
| 625 | 624 | { | |
| 626 | - throw new NotImplementedException(); | ||
| 627 | - //var fn = function_deserialization.setup_bare_concrete_function(proto, ) | ||
| 625 | + var fn = function_deserialization.recreate_function(proto, null); | ||
| 626 | + foreach (var name in proto.ConcreteFunctions) | ||
| 627 | + { | ||
| 628 | + _setup_function_captures(); | ||
| 629 | + } | ||
| 630 | + return (fn, setattr); | ||
| 628 | 631 | } | |
| 629 | 632 | ||
| 630 | 633 | private (ConcreteFunction, Action<object, object, object>) _recreate_bare_concrete_function(SavedBareConcreteFunction proto, | |
| Back | FazBrowse Home | New Git URL |
0 commit comments