Skip to content

Commit acae9b3

Browse files
committed
Partially support the analysis of loaded functions.
1 parent 4e6431e commit acae9b3

42 files changed

Lines changed: 782 additions & 284 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

TensorFlow.NET.sln

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11

22
Microsoft Visual Studio Solution File, Format Version 12.00
3-
# Visual Studio Version 16
4-
VisualStudioVersion = 16.0.31624.102
3+
# Visual Studio Version 17
4+
VisualStudioVersion = 17.4.33213.308
55
MinimumVisualStudioVersion = 10.0.40219.1
66
Project("{9A19103F-16F7-4668-BE54-9A1E7A4F7556}") = "Tensorflow.Binding", "src\TensorFlowNET.Core\Tensorflow.Binding.csproj", "{FD682AC0-7B2D-45D3-8B0D-C6D678B04144}"
77
EndProject
@@ -23,6 +23,8 @@ Project("{9A19103F-16F7-4668-BE54-9A1E7A4F7556}") = "Tensorflow.Keras.UnitTest",
2323
EndProject
2424
Project("{9A19103F-16F7-4668-BE54-9A1E7A4F7556}") = "TensorFlowNET.Graph.UnitTest", "test\TensorFlowNET.Graph.UnitTest\TensorFlowNET.Graph.UnitTest.csproj", "{3F5388FF-FBB4-462B-8F6F-829FFBAEB8A3}"
2525
EndProject
26+
Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "Tensorflow.Common", "Tensorflow.Common\Tensorflow.Common.csproj", "{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}"
27+
EndProject
2628
Global
2729
GlobalSection(SolutionConfigurationPlatforms) = preSolution
2830
Debug|Any CPU = Debug|Any CPU
@@ -153,6 +155,18 @@ Global
153155
{3F5388FF-FBB4-462B-8F6F-829FFBAEB8A3}.Release|x64.Build.0 = Release|x64
154156
{3F5388FF-FBB4-462B-8F6F-829FFBAEB8A3}.Release|x86.ActiveCfg = Release|Any CPU
155157
{3F5388FF-FBB4-462B-8F6F-829FFBAEB8A3}.Release|x86.Build.0 = Release|Any CPU
158+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Debug|Any CPU.ActiveCfg = Debug|Any CPU
159+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Debug|Any CPU.Build.0 = Debug|Any CPU
160+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Debug|x64.ActiveCfg = Debug|Any CPU
161+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Debug|x64.Build.0 = Debug|Any CPU
162+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Debug|x86.ActiveCfg = Debug|Any CPU
163+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Debug|x86.Build.0 = Debug|Any CPU
164+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Release|Any CPU.ActiveCfg = Release|Any CPU
165+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Release|Any CPU.Build.0 = Release|Any CPU
166+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Release|x64.ActiveCfg = Release|Any CPU
167+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Release|x64.Build.0 = Release|Any CPU
168+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Release|x86.ActiveCfg = Release|Any CPU
169+
{0C5DD8A8-AB1E-40AB-8CE3-F6EA0C1ED680}.Release|x86.Build.0 = Release|Any CPU
156170
EndGlobalSection
157171
GlobalSection(SolutionProperties) = preSolution
158172
HideSolutionNode = FALSE
Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
using OneOf;
2+
using System;
3+
4+
namespace Tensorflow.Common.Extensions
5+
{
6+
public static class OneofExtension
7+
{
8+
public static bool IsTypeOrDeriveFrom<T>(this IOneOf src)
9+
{
10+
return src.Value is T;
11+
}
12+
}
13+
}
Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,11 @@
1+
<Project Sdk="Microsoft.NET.Sdk">
2+
3+
<PropertyGroup>
4+
<TargetFramework>netstandard2.0</TargetFramework>
5+
</PropertyGroup>
6+
7+
<ItemGroup>
8+
<PackageReference Include="OneOf" Version="3.0.223" />
9+
</ItemGroup>
10+
11+
</Project>

src/TensorFlowNET.Core/Checkpoint/SaveUtil.cs

Lines changed: 11 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,12 @@
1-
using System;
1+
using OneOf;
2+
using System;
23
using System.Collections.Generic;
34
using System.Diagnostics;
45
using System.Linq;
56
using System.Text;
67
using Tensorflow.Train;
78
using Tensorflow.Training;
9+
using Tensorflow.Common.Extensions;
810
using pbc = global::Google.Protobuf.Collections;
911

1012
namespace Tensorflow.Checkpoint
@@ -28,7 +30,7 @@ Trackable object_to_save
2830
);
2931
public static class SaveUtil
3032
{
31-
public static (IDictionary<Trackable, IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>>>, IDictionary<Tensor, object>, IDictionary<string, IDictionary<string, Trackable>>, TrackableObjectGraph)
33+
public static (IDictionary<Trackable, IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>>>, IDictionary<Tensor, object>, IDictionary<string, IDictionary<string, Trackable>>, TrackableObjectGraph)
3234
serialize_graph_view(ObjectGraphView graph_view, IDictionary<Trackable, Trackable>? object_map = null, bool call_with_mapped_captures = false, object? cache = null)
3335
{
3436
var (trackable_data, node_ids) = gather_trackable_data(graph_view, object_map);
@@ -117,16 +119,16 @@ private static TrackableObjectGraph fill_object_graph_proto(IList<TrackableData>
117119
/// <param name="call_with_mapped_captures"></param>
118120
/// <param name="cache"></param>
119121
/// <param name="object_graph_proto"></param>
120-
private static IDictionary<Trackable, IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>>> get_and_write_tensors_to_serialize(IList<TrackableData> tensor_trackables, IDictionary<Trackable, int> node_ids,
122+
private static IDictionary<Trackable, IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>>> get_and_write_tensors_to_serialize(IList<TrackableData> tensor_trackables, IDictionary<Trackable, int> node_ids,
121123
bool call_with_mapped_captures, object? cache, TrackableObjectGraph object_graph_proto)
122124
{
123-
Dictionary<Trackable, IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>>> serialized_tensors = new();
125+
Dictionary<Trackable, IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>>> serialized_tensors = new();
124126
foreach(var td in tensor_trackables)
125127
{
126128
// TODO: deal with cache.
127129
var legacy_name = SaveableCompat.get_saveable_name(td.object_to_save) ?? "";
128130
Trackable trackable = null;
129-
IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>> tensor_dict;
131+
IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>> tensor_dict;
130132
if(!saveable_object_util.trackable_has_serialize_to_tensor(td.object_to_save) || legacy_name.Length > 0)
131133
{
132134
(trackable, tensor_dict) = get_tensors_from_legacy_saveable(td, node_ids, call_with_mapped_captures, object_graph_proto);
@@ -148,12 +150,12 @@ private static IDictionary<Trackable, IDictionary<string, Maybe<Tensor, IDiction
148150
return serialized_tensors;
149151
}
150152

151-
private static IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>> get_tensors_from_trackable(TrackableData trackable_data, bool call_with_mapped_captures, TrackableObjectGraph object_graph_proto)
153+
private static IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>> get_tensors_from_trackable(TrackableData trackable_data, bool call_with_mapped_captures, TrackableObjectGraph object_graph_proto)
152154
{
153155
var trackable = trackable_data.object_to_save;
154156

155157
// TODO: complete it. Note that actually `call_with_mapped_captures` is of function type.
156-
IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>> ret_tensor_dict;
158+
IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>> ret_tensor_dict;
157159
if (call_with_mapped_captures)
158160
{
159161
throw new NotImplementedException();
@@ -164,7 +166,7 @@ private static IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>> g
164166
}
165167

166168
// TODO: deal with the type `SaveSpce` (currently it will never be it).
167-
Dictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>> tensor_dict = new();
169+
Dictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>> tensor_dict = new();
168170
foreach(var pair in ret_tensor_dict)
169171
{
170172
var local_name = TrackableUtils.escape_local_name(pair.Key);
@@ -200,7 +202,7 @@ private static IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>> g
200202
/// <param name="call_with_mapped_captures"></param>
201203
/// <param name="object_graph_proto"></param>
202204
/// <returns></returns>
203-
private static (Trackable, IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>>) get_tensors_from_legacy_saveable(TrackableData trackable_data, IDictionary<Trackable, int> node_ids,
205+
private static (Trackable, IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>>) get_tensors_from_legacy_saveable(TrackableData trackable_data, IDictionary<Trackable, int> node_ids,
204206
bool call_with_mapped_captures, TrackableObjectGraph object_graph_proto)
205207
{
206208
Dictionary<Trackable, string> object_names = new();

src/TensorFlowNET.Core/Checkpoint/SaveUtilV1.cs

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
using pbc = global::Google.Protobuf.Collections;
99
using static Tensorflow.Binding;
1010
using Google.Protobuf;
11+
using OneOf;
1112

1213
namespace Tensorflow.Checkpoint;
1314

@@ -179,13 +180,13 @@ public static (IList<MySaveableObject>, object?) generate_saveable_objects(
179180

180181
// TODO: tensorflow python has a process with callable `saveable_factory`.
181182
List<MySaveableObject> saveables = new();
182-
if (maybe_saveable.TryGet<MySaveableObject>(out var s))
183+
if (maybe_saveable.TryPickT1(out var s, out var variable))
183184
{
184185
saveables.Add(s);
185186
}
186187
else
187188
{
188-
saveables.AddRange(saveable_object_util.saveable_objects_for_op(maybe_saveable.GetValue<BaseResourceVariable>() as Trackable, key));
189+
saveables.AddRange(saveable_object_util.saveable_objects_for_op(variable as Trackable, key));
189190
}
190191

191192
foreach (var saveable in saveables)
@@ -217,7 +218,7 @@ public static (IList<MySaveableObject>, object?) generate_saveable_objects(
217218

218219
public record class CheckpointFactoryData
219220
(
220-
Func<string, Maybe<BaseResourceVariable, MySaveableObject>> factory,
221+
Func<string, OneOf<BaseResourceVariable, MySaveableObject>> factory,
221222
string name,
222223
string checkpoint_key
223224
);

src/TensorFlowNET.Core/Checkpoint/checkpoint.cs

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
using Tensorflow.Operations;
1313
using Newtonsoft.Json;
1414
using Tensorflow.Training;
15+
using OneOf;
1516

1617
namespace Tensorflow.Checkpoint;
1718

@@ -49,7 +50,7 @@ public TrackableSaver(ObjectGraphView graph_view)
4950

5051
}
5152

52-
private (IDictionary<Trackable, IDictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>>>, IDictionary<Tensor, object>, IDictionary<string, IDictionary<string, Trackable>>, TrackableObjectGraph)
53+
private (IDictionary<Trackable, IDictionary<string, OneOf<Tensor, IDictionary<string, Tensor>>>>, IDictionary<Tensor, object>, IDictionary<string, IDictionary<string, Trackable>>, TrackableObjectGraph)
5354
gather_serialized_tensors(Tensor? object_graph_tensor = null)
5455
{
5556
var (serialized_tensors, feed_additions, registered_savers, graph_proto) = SaveUtil.serialize_graph_view(_graph_view, _object_map, cache:_cache);
@@ -68,7 +69,7 @@ public TrackableSaver(ObjectGraphView graph_view)
6869
Debug.Assert(!serialized_tensors.ContainsKey(Trackable.None) || !serialized_tensors[Trackable.None].ContainsKey(Trackable.Constants.OBJECT_GRAPH_PROTO_KEY));
6970
if (!serialized_tensors.ContainsKey(Trackable.None))
7071
{
71-
serialized_tensors[Trackable.None] = new Dictionary<string, Maybe<Tensor, IDictionary<string, Tensor>>>();
72+
serialized_tensors[Trackable.None] = new Dictionary<string, OneOf.OneOf<Tensor, IDictionary<string, Tensor>>>();
7273
}
7374
serialized_tensors[Trackable.None][Trackable.Constants.OBJECT_GRAPH_PROTO_KEY] = object_graph_tensor;
7475
return (serialized_tensors, feed_additions, registered_savers, graph_proto);
@@ -400,7 +401,7 @@ public void new_restore_ops(IEnumerable<Operation> new_ops)
400401
// skip the callback.
401402
}
402403

403-
public List<Operation> restore_saveables(Dictionary<string, Maybe<BaseResourceVariable, MySaveableObject>> tensor_saveables, List<CheckpointPosition> positions, object? registered_savers = null)
404+
public List<Operation> restore_saveables(Dictionary<string, OneOf<BaseResourceVariable, MySaveableObject>> tensor_saveables, List<CheckpointPosition> positions, object? registered_savers = null)
404405
{
405406
List<Operation> restore_ops = new();
406407
foreach(var position in positions)
@@ -412,7 +413,7 @@ public List<Operation> restore_saveables(Dictionary<string, Maybe<BaseResourceVa
412413
Dictionary<string, BaseResourceVariable> variable_dict = new();
413414
foreach(var item in tensor_saveables)
414415
{
415-
if(item.Value.TryGet<BaseResourceVariable>(out var variable))
416+
if(item.Value.TryPickT0(out var variable, out var _))
416417
{
417418
variable_dict[item.Key] = variable;
418419
}

0 commit comments

Comments
 (0)