Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ml_dataloader_TensorFlow.py
Go to the documentation of this file.
1### \file
2### \ingroup tutorial_ml
3### \notebook -nodraw
4### Example of getting batches of events from a ROOT dataset into a basic
5### TensorFlow workflow.
6###
7### \macro_code
8### \macro_output
9### \author Dante Niewenhuis
10
11import ROOT
12
13tree_name = "sig_tree"
14file_name = str(ROOT.gROOT.GetTutorialDir()) + "/machine_learning/data/Higgs_data.root"
15
16batch_size = 128
17approx_batches_in_memory = 50
18
19rdataframe = ROOT.RDataFrame(tree_name, file_name)
20target = ["Type"]
21
22# Returns two TF.Dataset for training and validation batches.
24 rdataframe,
25 batch_size,
26 approx_batches_in_memory,
27 target=target,
28 shuffle=True,
29 drop_remainder=True,
30)
31
32ds_train, ds_valid = dl.train_test_split(test_size=0.3)
33
34num_of_epochs = 2
35
36# Datasets have to be repeated as many times as there are epochs
37ds_train_repeated = ds_train.as_tensorflow().repeat(num_of_epochs)
38ds_valid_repeated = ds_valid.as_tensorflow().repeat(num_of_epochs)
39
40# Number of batches per epoch must be given for model.fit
41train_batches_per_epoch = ds_train.num_batches
42validation_batches_per_epoch = ds_valid.num_batches
43
44# Get a list of the columns used for training
45input_columns = ds_train.feature_columns
46num_features = len(input_columns)
47
48##############################################################################
49# AI example
50##############################################################################
51# TensorFlow has to be imported after ROOT has initialized its interpreter, to avoid
52# LLVM symbol clashes with TensorFlow>=2.20.0.
53import tensorflow as tf # noqa: E402
54
55# Define TensorFlow model
57 [
58 tf.keras.layers.Input(shape=(num_features,)),
59 tf.keras.layers.Dense(300, activation=tf.nn.tanh),
60 tf.keras.layers.Dense(300, activation=tf.nn.tanh),
61 tf.keras.layers.Dense(300, activation=tf.nn.tanh),
63 ]
64)
65
67model.compile(optimizer="adam", loss=loss_fn, metrics=["accuracy"])
68
70 ds_train_repeated,
71 steps_per_epoch=train_batches_per_epoch,
72 validation_data=ds_valid_repeated,
73 validation_steps=validation_batches_per_epoch,
74 epochs=num_of_epochs,
75)
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h Atom_t Int_t ULong_t ULong_t unsigned char prop_list Atom_t Atom_t Atom_t Time_t UChar_t len
ROOT's RDataFrame offers a modern, high-level interface for analysis of data stored in TTree ,...