Generative Adversarial Network Action Set

Hmeq Data

This section contains PROC CAS code.

Note: Input data must be accessible in your CAS session, either as a CAS table or as a transient-scope table. A CAS table has a two-level name: the first level is your CAS engine libref, and the second level is the table name. You refer to this table in the CAS procedure by specifying only the second level. For more information about two-level names, see Chapter 2, Shared Concepts (SAS Visual Data Mining and Machine Learning: Procedures). A transient-scope table is called directly from the action and exists in memory for the duration of the action. For more information about accessing data, see SAS Viya: System Programming Guide. For more information about PROC CAS and programming in CASL, see SAS Cloud Analytic Services: CASL Programmer’s Guide and SAS Cloud Analytic Services: CASL Reference.

This example shows how to use the tabularGanTrain action to analyze and augment the hmeq data table, which contains data about mortgage applicants. The following DATA step creates the input data table mycas.hmeq in your CAS session. These statements assume that your CAS engine libref is named mycas, but you can substitute any appropriately defined CAS engine libref.

data mycas.hmeq;
   set sampsio.hmeq;
run;

The following code uses the GMM action to cluster the value and clage variables and to generate the centroids table to be used in the tabularGanTrain action. This step saves the weight, centroid ID, mean, and standard deviation for each cluster in the centroids data table.

 /* Define the runloop macro */
 %macro runloop;
     /* Specify all interval input variables*/
     %let names=value clage;
     /* Loop over all variables that need centroids generation */
     %do i=1 %to %sysfunc(countw(&names));
         %let name&i = %scan(&names, &i, %str( ));
         /* Call the GMM action to cluster each variable */
         proc cas ;
             action nonParametricBayes.gmm result=R/
                 table       = {name="hmeq"},
                 inputs      = {"&&name&i"},/*'value'*/
                 seed        = 1234567890,
                 maxClusters = 10,
                 alpha       = 1,
                 infer       = {method="VB",
                                maxVbIter =30,
                                covariance="diagonal",
                                threshold=0.01},
                 output      = {casOut={name='Score', replace=true},
                                copyVars={'value'}},
                 display     = {names={ "ClusterInfo"}}
                ;
             run;
             saveresult R.ClusterInfo replace dataset=work.weights&i;
         run;
         quit;

         /* Save variable name, weights, mean,     */
         /* and standard deviation of each cluster */
         data  weights&i;
             varname = "&&name&i";
             set  weights&i(rename=(&&name&i.._Mean=Mean
                                    &&name&i.._Variance=Var));
             /* Calculate standard deviation from variance*/
             std = sqrt(Var);
             drop Var;
         run;

         /* Construct centroids table from saved weights */
         %if &i=1 %then %do;
             data centroids;
             set weights&i;
             run;
         %end;
         %else %do;
             data centroids;
             set centroids weights&i;
             run;
         %end;
     %end;
 %mend;

 /* Run the runloop macro to generate the centroids table */
 %runloop;

The following DATA step uploads the centroids table to your CAS session:

data mycas.centroids;
   set centroids;
run;

The following PROC CAS statements use the tabularGanTrain action to train a tabular GAN model and save the trained model in an analytic store. The PROC PRINT statement prints the generated samples.


 proc cas;
     loadactionset "generativeAdversarialNet";
     action tabularGanTrain result = r /
         table           = {name = "hmeq",
                            vars = {'bad','value','clage','job'}},
         centroidsTable  = "centroids",
         nominals        = {"bad","job"},
         gpu             = {useGPU = True, device = 0},
         optimizerAe     = {method = "ADAM", numEpochs = 3},
         optimizerGan    = {method = "ADAM", numEpochs = 5},
         seed            = 12345,
         scoreSeed       = 0,
         numSamples      = 5,
         saveState       = {name = 'cpctStore', replace = True},
         casOut          = {name = 'out', replace = True};
     print r;
 run;
 quit;

 proc print data = mycas.out;
 run;

The table parameter names the input data table, hmeq, to be analyzed. The vars subparameter specifies the input variables in the hmeq, including bad, value, clage, and job. The centroidsTable parameter specifies the centroids column table. The nominals parameter specifies the nominal variables from the input. The method subparameter in the optimizerAe parameter list specifies that the autoencoder optimizer be used in the training process; currently only the Adam method is supported. The numEpochs subparameter in the optimizerAe parameter list specifies the number of epochs to use for training the autoencoder; one epoch processes all the training data once. The method subparameter in the optimizerGAN parameter list specifies the GAN optimizer to use in the training process; currently only the Adam method is supported. The numEpochs subparameter in the optimizerGAN parameter list specifies the number of epochs to use for training the GAN model; one epoch processes all the training data once. The seed parameter specifies the random seed 12345 be used in the training process. The numSamples parameter specifies the number of synthetic samples to generate. The useGPU subparameter in the gpu parameter list specifies whether you want to run the action on the GPU. If you set it to True, then the action uses the GPU device whose ID you specify in the device subparameter. The scoreSeed parameter specifies that the random seed 0 be used in generating the out data table. The saveState parameter saves the trained model to the cpctStore data table in the active caslib for future scoring. The casOut parameter outputs the score output to the out data table.

The tabularGanTrain action generates four ODS tables, which are shown in Output 17.2.1 through Output 17.2.4.

Output 17.2.1: Iteration History

r: Results from generativeAdversarialNet.tabularGanTrain

Iteration History
Epoch NumberAutoencoder LossGenerator LossDiscriminator Loss
10.0794085786..
20.0261880551..
30.0126868188..
1.1.2421702147-0.010049699
2.1.25833487510.0065689385
3.1.2131823301-0.006526524
4.1.1670466661-0.005588347
5.1.2273106575-0.020535301


Output 17.2.2: Number of Observations

Number of Observations
Number of Observations Read5960
Number of Observations Used5440


Output 17.2.3: Model Information

Model Information
Generator Embedding Dimension128
Number of Observations in One Minibatch500
Number of Observations Group Together in Applying the Discriminator10
Weight for Regularizing the Discriminator10
Exponential Decay Rate for the First-Moment Estimates for the Autoencoder's Optimizer0.900000
Exponential Decay Rate for the Second-Moment Estimates for the Autoencoder's Optimizer0.999000
Learning Rate for the Autoencoder's Optimizer0.001
Number of Epochs for the Autoencoder's Training3
Weight Decay for the Autoencoder's Optimizer1e-08
Exponential Decay Rate for the First-Moment Estimates to Optimize GAN0.500000
Exponential Decay Rate for the Second-Moment Estimates to Optimize GAN0.900000
Learning Rate for the GAN Optimizer2e-05
Number of Epochs for the GAN Training5
Weight Decay for the Generator's Optimizer0.0001
Weight Decay for the Discriminator's Optimizer1e-06
Seed for Random Initialization12345
Whether to Use Log Frequency of Categorical Levels in the Conditional SamplingTrue


Output 17.2.4: Level Frequencies for Nominal Variables

Level Frequencies for Nominal
Variables
Variable
Name
LevelFrequency
BAD04430
BAD11010
JOBMgr730
JOBOffice925
JOBOther2265
JOBProfExe1228
JOBSales108
JOBSelf184


Output 17.2.5: CAS Output

ObsBADVALUECLAGEJOB
11173650.4629.837Mgr
2028091.3463.892Other
3166033.3263.206ProfExe
41-414108.69312.082Mgr
51206353.02161.891Office


In this example, the tabularGanTrain action generates the score output table out in the active caslib. The out table contains five generated samples from the trained model.

Hmeq Data

This section contains Python code for the analysis in the CASL version of this example, which contains details about the results.

Note: In order to run this code, the data that are described in the CASL version need to be accessible to the CAS server. One way to do this is to convert the hmeq data to the comma-separated-value (CSV) file hmeq.csv and then use the following code to load the CSV file into CAS:

s.upload_file('hmeq.csv')

For more information about coding in Python, see Getting Started with SAS Viya for Python and SAS Viya: System Programming Guide.

The following Python code loads the generativeAdversarialNet action set and then uses the tabularGanTrain action to train a tabular model on the hmeq data:



 # read csv data and upload to CAS
 s.upload("hmeq.csv", casout=dict(name='hmeq', replace=True)).casTable

 # print all columns
 print(hmeq_data.columns)

 # define the continuous columns to generate the centroids table
 # here "VALUE" and "CLAGE" are the continuous columns
 ContCols = ['VALUE', 'CLAGE']


 # load gmm action set and calculate the centroids table
 s.loadactionset('nonParametricBayes')

 Cent = pd.DataFrame()
 for col in ContCols:

     s.gmm(
        table       = {"name":"hmeq"},
                 inputs      = {col},
                 seed        = 1234567890,
                 maxClusters = 10,
                 alpha       = 1,
                 infer       = {"method":"VB",
                                "maxVbIter" :30,
                                "covariance":"diagonal",
                                "threshold":0.01},
                 output={"casOut":{"name":"Score", "replace":"True"}},
                 display     = {"names":{ "ClusterInfo"}},
                 outputtables ={"names":{"ClusterInfo"}},
     )
     cluster_value=s.CASTable("ClusterInfo",replace="True" )
     cluster_value=cluster_value.to_frame()
     table_value=cluster_value.rename(columns={col+'_Mean':'Mean'})
     table_value=table_value.rename(columns={col+'_Variance':'Variance'})
     table_value['Std']=table_value['Variance']**1/2
     table_value=table_value.drop(columns={'Variance'})
     table_value.insert(0,'VarName',col)
     table_value=table_value.rename(columns={'_CLUSTER_ID_':'Centroid-i'})
     s.table.dropTable("ClusterInfo")
     Cent = Cent.append(table_value)

 Cent.rename(columns={'Cweight':'Weight'}, inplace=True)
 s.upload(Cent, casout=dict(name='cen_table', replace=True)).casTable


 # load generativeAdversarialNet action set
 s.loadactionset('generativeAdversarialNet')

 results = s.tabularGanTrain(
 table = {"name":"hmeq","vars" :{"value","clage","bad","job"}},
     centroidsTable= "cen_table",
     gpu = {"useGPU":True},
     nominals ={"bad","job"},
     optimizerAe ={"method":'ADAM',"numEpochs":5},
     optimizerGan ={"method":'ADAM',"numEpochs":5},
     seed = 12345,
     scoreSeed = 1234,
     numSamples =10,
     saveState ={"name":"cpctStore", "replace":True},
     casOut = {"name":"out", "replace":True}
 )
 s.table.fetch(table={"name":"out"})


For details about the results of this analysis, see the CASL version of this example.

Hmeq Data

This example is not available for the Lua programming language.

Hmeq Data

This example is not available for the R programming language.

Last updated: September 15, 2022