modelflowib 2.73__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
modeldekom.py ADDED
@@ -0,0 +1,651 @@
1
+ # -*- coding: utf-8 -*-
2
+ """
3
+ Module for making attribution analysis of a model.
4
+
5
+ The main function is attribution
6
+
7
+ Created on Wed May 31 08:50:51 2017
8
+
9
+ @author: hanseni
10
+
11
+ """
12
+
13
+
14
+ import pandas as pd
15
+ import fnmatch
16
+ import matplotlib.pyplot as plt
17
+ import matplotlib as mpl
18
+ import matplotlib.dates as mdates
19
+
20
+ import numpy
21
+ import ipywidgets as ip
22
+ import pdb
23
+
24
+
25
+
26
+ from modelhelp import cutout
27
+ import modelclass as mc
28
+ #import modeldekom as mk
29
+ import modelvis as mv
30
+
31
+
32
+ idx= pd.IndexSlice
33
+
34
+ def attribution(model,experiments,start='',end='',save='',maxexp=10000,showtime=False,
35
+ summaryvar=['*']
36
+ ,silent=False,msilent=True,type='level'):
37
+ """ Calculates an attribution analysis on a model
38
+ accepts a dictionary with experiments. the key is experiment name, the value is a list
39
+ of variables which has to be reset to the values in the baseline dataframe. """
40
+ summaryout = model.vlist(summaryvar)
41
+ adverseny = model.lastdf
42
+ base = model.basedf
43
+ if type == 'level':
44
+ adverse0=adverseny[summaryout].loc[start:end,:].copy()
45
+ elif type == 'growth':
46
+ adverse0=adverseny[summaryout].pct_change().loc[start:end,:].copy() * 100.
47
+ ret={}
48
+ modelsave = model.save # save the state of model.save
49
+ model.save = False # no need to save the experiments in each run
50
+ with model.timer('Total dekomp',showtime):
51
+ for i,(e,var) in enumerate(experiments.items()):
52
+ if i >= maxexp : break # when we are testing
53
+ oldvar=adverseny[var].copy()
54
+ if not silent:
55
+ print(i,'Experiment :',e,'\n','Touching: \n', var)
56
+ adverseny[var] = base[var]
57
+ tempdf = model(adverseny ,start,end,samedata=True,
58
+ silent=msilent)[summaryout]
59
+ adverseny[var] = oldvar
60
+
61
+ if type == 'level':
62
+ ret[e] = tempdf[summaryout].loc[start:end,:]
63
+ elif type == 'growth':
64
+ ret[e] = tempdf.pct_change().loc[start:end,:] * 100.
65
+
66
+ difret = {e : adverse0-ret[e] for e in ret}
67
+
68
+ df = pd.concat([difret[v] for v in difret],keys=difret.keys()).T
69
+ if save:
70
+ df.to_pickle('data\\' +save +r'.pc')
71
+
72
+ model.save = modelsave # restore the state of model.save
73
+ return df
74
+
75
+ def attribution_new(model,experiments,start='',end='',save='',maxexp=10000,showtime=False,
76
+ summaryvar=['*']
77
+ ,silent=False,msilent=True,type='level'):
78
+ """
79
+ Performs attribution analysis on a model and returns the decomposition of differences
80
+ between baseline and experimental scenarios in terms of 'level' or 'growth'.
81
+
82
+ This function calculates the impact of resetting specified variables to their baseline
83
+ values across multiple experiments and returns the resulting differences.
84
+
85
+ Parameters:
86
+ model (object): The model instance to perform attribution analysis on.
87
+ It should support methods like `vlist`, `lastdf`, and `basedf`.
88
+ experiments (dict): A dictionary where keys are experiment names, and values
89
+ are lists of variables to reset to baseline values.
90
+ start (str, optional): The start period for the analysis. Defaults to an empty string.
91
+ end (str, optional): The end period for the analysis. Defaults to an empty string.
92
+ save (str, optional): If provided, saves the resulting data to a file with this name
93
+ (appends '_level' and '_growth' to distinguish the types). Defaults to an empty string.
94
+ maxexp (int, optional): Maximum number of experiments to process. Defaults to 10,000.
95
+ showtime (bool, optional): Whether to display timing information for the operation. Defaults to `False`.
96
+ summaryvar (list, optional): List of variables to include in the analysis output.
97
+ Defaults to `['*']`, which includes all variables.
98
+ silent (bool, optional): If `True`, suppresses print statements. Defaults to `False`.
99
+ msilent (bool, optional): If `True`, suppresses print statements from the model. Defaults to `True`.
100
+ type (str, optional): The type of analysis to perform, either 'level' or 'growth'. Defaults to 'level'.
101
+
102
+ Returns:
103
+ dict: A dictionary containing two DataFrames:
104
+ - 'level': Decomposition of differences in levels for the specified variables.
105
+ - 'growth': Decomposition of differences in growth rates (percentage change).
106
+
107
+ Example:
108
+ # Example usage with a model and experiments
109
+ results = attribution_new(
110
+ model=my_model,
111
+ experiments={
112
+ 'experiment1': ['variable1', 'variable2'],
113
+ 'experiment2': ['variable3']
114
+ },
115
+ start='2020Q1',
116
+ end='2021Q4',
117
+ summaryvar=['variable_summary'],
118
+ type='growth'
119
+ )
120
+
121
+ # Access the level results
122
+ level_results = results['level']
123
+
124
+ # Access the growth results
125
+ growth_results = results['growth']
126
+ """
127
+ summaryout = model.vlist(summaryvar)
128
+ adverseny = model.lastdf
129
+ base = model.basedf
130
+ adverse0_level = adverseny[summaryout].loc[start:end,:].copy()
131
+ adverse0_growth = adverseny[summaryout].pct_change().loc[start:end,:].copy() * 100.
132
+ ret_level = {}
133
+ ret_growth = {}
134
+ modelsave = model.save # save the state of model.save
135
+ model.save = False # no need to save the experiments in each run
136
+ with model.timer('Total dekomp',showtime):
137
+ for i,(e,var) in enumerate(experiments.items()):
138
+ if i >= maxexp : break # when we are testing
139
+ oldvar=adverseny[var].copy()
140
+ if not silent:
141
+ print(i,'Experiment :',e,'\n','Touching: \n', var)
142
+ adverseny[var] = base[var]
143
+ tempdf = model(adverseny ,start,end,samedata=True,
144
+ silent=msilent)[summaryout]
145
+ adverseny[var] = oldvar
146
+
147
+ ret_level[e] = tempdf.loc[start:end,:]
148
+ ret_growth[e] = tempdf.pct_change().loc[start:end,:] * 100.
149
+
150
+ difret_level = {e : adverse0_level - ret_level[e] for e in ret_level}
151
+ difret_growth = {e : adverse0_growth - ret_growth[e] for e in ret_growth}
152
+
153
+ df_level = pd.concat([difret_level[v] for v in difret_level],keys=difret_level.keys()).T
154
+ df_growth = pd.concat([difret_growth[v] for v in difret_growth],keys=difret_growth.keys()).T
155
+ if save:
156
+ df_level.to_pickle('data\\' +save +r'_level.pc')
157
+ df_level.to_pickle('data\\' +save +r'_growth.pc')
158
+
159
+ model.save = modelsave # restore the state of model.save
160
+ return {'level':df_level, 'growth':df_growth}
161
+
162
+
163
+ def ilist(df,pat):
164
+ '''returns a list of variable in the model matching the pattern,
165
+ the pattern can be a list of patterns of a sting with patterns seperated by
166
+ blanks
167
+
168
+ This function operates on the index names of a dataframe. Relevant for attribution analysis
169
+ '''
170
+ if isinstance(pat,list):
171
+ upat=pat
172
+ else:
173
+ upat = [pat]
174
+
175
+ ipat = upat
176
+ out = [v for p in ipat for up in p.split() for v in sorted(fnmatch.filter(df.index,up.upper()))]
177
+ return out
178
+
179
+
180
+ def GetAllImpact(impact,sumaryvar):
181
+ ''' get all the impact from at impact dataframe'''
182
+ exo = list({v for v,t in impact.columns})
183
+ df = pd.concat([impact.loc[sumaryvar,c] for c in exo],axis=1)
184
+ df.columns = exo
185
+ return df
186
+
187
+ def GetSumImpact(impact,pat='PD__*'):
188
+ """Gets the accumulated differences attributet to each impact group """
189
+ a = impact.loc[ilist(impact,pat),:].T.groupby(level=0).sum().T
190
+ return a
191
+
192
+ def GetLastImpact(impact,pat='RCET1__*'):
193
+ """Gets the last differences attributet to each impact group """
194
+ # assert 1==2
195
+ a = impact.loc[ilist(impact,pat),:].T.groupby(level=0).last().T
196
+ return a
197
+
198
+ def GetAllImpact(impact,pat='RCET1__*'):
199
+ """Gets the last differences attributet to each impact group """
200
+ a = impact.loc[ilist(impact,pat),:]
201
+ return a
202
+
203
+ def _as_title(desc):
204
+ """Coerce a description (which may be a list/tuple) into a plain string,
205
+ so it can be passed as a matplotlib/pandas plot title."""
206
+ if isinstance(desc, (list, tuple)):
207
+ return ' '.join(str(d) for d in desc)
208
+ return desc
209
+
210
+ def GetOneImpact(impact,pat='RCET1__*',per=''):
211
+ """Gets differences attributet to each impact group in period:per """
212
+ a = impact.loc[ilist(impact,pat),idx[:,per]]
213
+ a.columns = [v[0] for v in a.columns]
214
+ return a
215
+
216
+ def AggImpact(impact):
217
+ """ Calculates the sum of impacts and place in the last column
218
+
219
+ This function is applied to the result iof a Get* function"""
220
+ asum= impact.sum(axis=1)
221
+ asum.name = '_Sum'
222
+ aout = pd.concat([impact,asum],axis=1)
223
+ return aout
224
+
225
+
226
+ class totdif():
227
+ """
228
+ A class for performing model-wide attribution analysis.
229
+
230
+ This class is designed to analyze and attribute differences in model outputs
231
+ based on experiments that alter specified variables. It provides methods
232
+ to decompose, visualize, and explain the impacts of different variable changes
233
+ on model outputs.
234
+
235
+ Parameters:
236
+ model (object): The model instance to perform attribution analysis on.
237
+ The model should implement methods like `exodif`, `current_per`, and `vlist`.
238
+ summaryvar (str or list, optional): Variables to summarize in the analysis.
239
+ Defaults to '*', which includes all variables.
240
+ desdic (dict, optional): A dictionary mapping variable names to descriptions,
241
+ used for labeling and visualization. Defaults to an empty dictionary.
242
+ experiments (dict, optional): A dictionary where keys are experiment names
243
+ and values are lists of variables to be reset to baseline values.
244
+ Defaults to None, which uses all variables with differences.
245
+
246
+ Attributes:
247
+ diffdf (pd.DataFrame): DataFrame containing differences between baseline and adverse scenarios.
248
+ diffvar (Index): List of variables with differences, obtained from the model's `exodif` method.
249
+ go (bool): Flag indicating whether there are variables to attribute to.
250
+ experiments (dict): Dictionary mapping variable names to experiment names. Defaults to all variables with differences.
251
+ model: Reference to the model being analyzed.
252
+ start (str): Start period of the analysis.
253
+ end (str): End period of the analysis.
254
+ desdic (dict): Dictionary for describing variables in visualizations.
255
+ summaryvar (str or list): Variables to summarize in the output.
256
+ summaryout: Processed list of summary variables obtained from the model's `vlist` method.
257
+ res (dict): Results of the attribution analysis for both 'level' and 'growth' perspectives.
258
+
259
+ Methods:
260
+ explain_last(pat='', top=0.9, title='', use='level', threshold=0.0, ysize=5):
261
+ Visualizes the decomposition for the last period.
262
+
263
+ explain_sum(pat='', top=0.9, title='', use='level', threshold=0.0, ysize=5):
264
+ Visualizes the summed decomposition over all periods.
265
+
266
+ explain_per(pat='', per='', top=0.9, title='', use='level', threshold=0.0, ysize=5):
267
+ Visualizes the decomposition for a specific period.
268
+
269
+ explain_all(pat='', stacked=True, kind='bar', top=0.9, title='', use='level',
270
+ threshold=0.0, resample='', axvline=None):
271
+ Visualizes the decomposition for all periods with optional customization.
272
+
273
+ totexplain(pat='*', vtype='all', stacked=True, kind='bar', per='', top=0.9, title='',
274
+ use='level', threshold=0.0, ysize=10, **kwargs):
275
+ A wrapper method that selects the appropriate visualization based on the type of data to attribute.
276
+
277
+ Example:
278
+ # Initialize the `totdif` class with a model and run attribution analysis
279
+ td = totdif(model=my_model, summaryvar='*', desdic=description_dict)
280
+
281
+ # Visualize the last period's decomposition
282
+ fig = td.explain_last(pat='variable_pattern', use='level')
283
+ fig.show()
284
+
285
+ # Summed decomposition visualization
286
+ fig = td.explain_sum(top=0.8, title='Cumulative Impact')
287
+ fig.show()
288
+ """
289
+
290
+ def __init__(self, model,summaryvar='*',desdic={},experiments = None):
291
+
292
+ self.diffdf = model.exodif()
293
+ self.diffvar = self.diffdf.columns
294
+ if len(self.diffvar) == 0:
295
+ print('No variables to attribute to ')
296
+ self.go = False
297
+ self.typetext = 'Unknown'
298
+
299
+ else:
300
+ self.go = True
301
+ self.experiments = {v:v for v in self.diffvar} if experiments == None else experiments
302
+ self.model = model
303
+ self.start = self.model.current_per.tolist()[0]
304
+ self.end = self.model.current_per.tolist()[-1]
305
+
306
+ self.desdic = desdic
307
+ self.summaryvar = summaryvar
308
+ self.summaryout = model.vlist(self.summaryvar)
309
+
310
+ self.res = attribution_new(self.model,self.experiments,self.start,self.end,
311
+ summaryvar=self.summaryvar,showtime=1,silent=1,type=type)
312
+
313
+ def explain_last(self,pat='',top=0.9,title='',use='level',threshold=0.0,ysize=5):
314
+ '''
315
+ Explains last period
316
+
317
+ Args:
318
+ pat (TYPE, optional): DESCRIPTION. Defaults to ''.
319
+ top (TYPE, optional): DESCRIPTION. Defaults to 0.9.
320
+ title (TYPE, optional): DESCRIPTION. Defaults to ''.
321
+ use (TYPE, optional): DESCRIPTION. Defaults to 'level'.
322
+ threshold (TYPE, optional): DESCRIPTION. Defaults to 0.0.
323
+ ysize (TYPE, optional): DESCRIPTION. Defaults to 5.
324
+
325
+ Returns:
326
+ fig (TYPE): DESCRIPTION.
327
+
328
+ '''
329
+ # assert 1==2
330
+ if self.go:
331
+ self.impact = GetLastImpact(self.res[use],pat=pat).T.rename(index=self.desdic)
332
+ ntitle = f'Decomposition last period, {use}' if title == '' else title
333
+ fig = mv.waterplot(self.impact,autosum=1,allsort=1,top=top,title= ntitle,desdic=self.desdic,
334
+ threshold=threshold,ysize=ysize)
335
+ return fig
336
+
337
+ def explain_sum(self,pat='',top=0.9,title='',use='level',threshold=0.0,ysize=5):
338
+ '''
339
+ Explains the sum
340
+
341
+
342
+ Args:
343
+ pat (TYPE, optional): DESCRIPTION. Defaults to ''.
344
+ top (TYPE, optional): DESCRIPTION. Defaults to 0.9.
345
+ title (TYPE, optional): DESCRIPTION. Defaults to ''.
346
+ use (TYPE, optional): DESCRIPTION. Defaults to 'level'.
347
+ threshold (TYPE, optional): DESCRIPTION. Defaults to 0.0.
348
+ ysize (TYPE, optional): DESCRIPTION. Defaults to 5.
349
+
350
+ Returns:
351
+ fig (TYPE): DESCRIPTION.
352
+
353
+ '''
354
+ if self.go:
355
+ self.impact = GetSumImpact(self.res[use],pat=pat).T.rename(index=self.desdic)
356
+ ntitle = f'Decomposition, sum over all periods, {use}' if title == '' else title
357
+ fig = mv.waterplot(self.impact,autosum=1,allsort=1,top=top,title=ntitle,desdic=self.desdic,
358
+ threshold=threshold,ysize=ysize )
359
+ return fig
360
+
361
+ def explain_per(self,pat='',per='',top=0.9,title='',use='level',threshold=0.0,ysize=5):
362
+ '''
363
+ Explains a periode
364
+
365
+ Args:
366
+ pat (TYPE, optional): DESCRIPTION. Defaults to ''.
367
+ per (TYPE, optional): DESCRIPTION. Defaults to ''.
368
+ top (TYPE, optional): DESCRIPTION. Defaults to 0.9.
369
+ title (TYPE, optional): DESCRIPTION. Defaults to ''.
370
+ use (TYPE, optional): DESCRIPTION. Defaults to 'level'.
371
+ threshold (TYPE, optional): DESCRIPTION. Defaults to 0.0.
372
+ ysize (TYPE, optional): DESCRIPTION. Defaults to 5.
373
+
374
+ Returns:
375
+ fig (TYPE): DESCRIPTION.
376
+
377
+ '''
378
+ if self.go:
379
+ tper = self.res[use].columns.get_level_values(1)[0] if per == '' else per
380
+ self.impact = GetOneImpact(self.res[use],pat=pat,per=tper).T.rename(index=self.desdic)
381
+ t2per = str(tper.date()) if type(tper) == pd._libs.tslibs.timestamps.Timestamp else tper
382
+ ntitle = f'Decomposition, {use}: {t2per}' if title == '' else title
383
+ fig = mv.waterplot(self.impact,autosum=1,allsort=1,top=top,title=ntitle,desdic=self.desdic ,
384
+ threshold=threshold,ysize=ysize)
385
+ return fig
386
+
387
+
388
+ def explain_allold(self,pat='',stacked=True,kind='bar',top=0.9,title='',use='level',
389
+ threshold=0.0,resample='',axvline=None):
390
+ if self.go:
391
+ years = mdates.YearLocator() # every year
392
+ months = mdates.MonthLocator() # every month
393
+ years_fmt = mdates.DateFormatter('%Y')
394
+
395
+ selected = GetAllImpact(self.res[use],pat)
396
+ grouped = selected.stack().groupby(level=0)
397
+ fig, axis = plt.subplots(nrows=len(grouped),ncols=1,figsize=(10,5*len(grouped)),constrained_layout=False)
398
+ width = 0.5 # the width of the barsser
399
+ ntitle = f'Decomposition, {use}' if title == '' else title
400
+ laxis = axis if isinstance(axis,numpy.ndarray) else [axis]
401
+ for j,((name,dfatt),ax) in enumerate(zip(grouped,laxis)):
402
+ dfatt.index = [i[1] for i in dfatt.index]
403
+ if resample=='':
404
+ tempdf=cutout(dfatt.T,threshold).T
405
+ else:
406
+ tempdf=cutout(dfatt.T,threshold).T.resample(resample).mean()
407
+ # pdb.set_trace()
408
+ tempdf.plot(ax=ax,kind=kind,stacked=stacked,title=_as_title(self.desdic.get(name,name)))
409
+ ax.set_ylabel(name,fontsize='x-large')
410
+ # ax.set_xticklabels(tempdf.index.tolist(), rotation = 45,fontsize='x-large')
411
+ ## ax.xaxis.set_minor_locator(plt.NullLocator())
412
+ ## ax.tick_params(axis='x', labelleft=True)
413
+ # ax.xaxis.set_major_locator(years)
414
+ # ax.xaxis_date()
415
+ # ax.xaxis.set_major_formatter(years_fmt)
416
+ # ax.xaxis.set_minor_locator(months)
417
+ # ax.tick_params(axis='x', labelrotation=45,right = True)
418
+ if type(axvline) != type(None): axis.axvline(axvline)
419
+ fig.suptitle(ntitle,fontsize=20)
420
+ if 1:
421
+ # plt.tight_layout()
422
+ # fig.subplots_adjust(top=top)
423
+ fig.set_constrained_layout(True)
424
+
425
+ return fig
426
+
427
+ def explain_all(self,pat='',stacked=True,kind='bar',top=0.9,title='',use='level',
428
+ threshold=0.0,resample='',axvline=None):
429
+ '''
430
+ Explains all
431
+
432
+ Args:
433
+ pat (TYPE, optional): DESCRIPTION. Defaults to ''.
434
+ stacked (TYPE, optional): DESCRIPTION. Defaults to True.
435
+ kind (TYPE, optional): DESCRIPTION. Defaults to 'bar'.
436
+ top (TYPE, optional): DESCRIPTION. Defaults to 0.9.
437
+ title (TYPE, optional): DESCRIPTION. Defaults to ''.
438
+ use (TYPE, optional): DESCRIPTION. Defaults to 'level'.
439
+ threshold (TYPE, optional): DESCRIPTION. Defaults to 0.0.
440
+ resample (TYPE, optional): DESCRIPTION. Defaults to ''.
441
+ axvline (TYPE, optional): DESCRIPTION. Defaults to None.
442
+
443
+ Returns:
444
+ None.
445
+
446
+ '''
447
+ import warnings
448
+ if self.go:
449
+ years = mdates.YearLocator() # every year
450
+ months = mdates.MonthLocator() # every month
451
+ years_fmt = mdates.DateFormatter('%Y')
452
+
453
+ selected = GetAllImpact(self.res[use],pat)
454
+ with warnings.catch_warnings():
455
+ warnings.simplefilter('ignore', FutureWarning)
456
+
457
+ grouped = selected.stack().groupby(level=0)
458
+ fig, axis = plt.subplots(nrows=len(grouped),ncols=1,figsize=(10,5*len(grouped)),constrained_layout=False)
459
+ width = 0.5 # the width of the barsser
460
+ ntitle = f'Decomposition, {use}' if title == '' else title
461
+ laxis = axis if isinstance(axis,numpy.ndarray) else [axis]
462
+ with warnings.catch_warnings():
463
+ warnings.simplefilter('ignore', FutureWarning)
464
+ for j,((name,dfatt),ax) in enumerate(zip(grouped,laxis)):
465
+ dfatt.index = [i[1] for i in dfatt.index]
466
+ if resample=='':
467
+ tempdf=cutout(dfatt.T,threshold).T
468
+ else:
469
+ tempdf=cutout(dfatt.T,threshold).T.resample(resample).mean()
470
+ # pdb.set_trace()
471
+ selfstack = (kind == 'line' or kind == 'area') and stacked
472
+ tempdf = tempdf.rename(columns=self.desdic)
473
+ with warnings.catch_warnings():
474
+ warnings.simplefilter("ignore", category=UserWarning)
475
+ if selfstack:
476
+ df_neg, df_pos =tempdf.clip(upper=0), tempdf.clip(lower=0)
477
+ df_pos.plot(ax=ax,kind=kind,stacked=stacked,title=_as_title(self.desdic.get(name,name)))
478
+ ax.set_prop_cycle(None)
479
+ df_neg.plot(ax=ax,legend=False,kind=kind,stacked=stacked,title=_as_title(self.desdic.get(name,name)))
480
+ ax.set_ylim([df_neg.sum(axis=1).min(), df_pos.sum(axis=1).max()])
481
+ else:
482
+ tempdf.plot(ax=ax,kind=kind,stacked=stacked,title=_as_title(self.desdic.get(name,name)))
483
+ if len(tempdf.index) < 9:
484
+ ax.set_xticks(range(len(tempdf.index)))
485
+ ax.set_xticklabels(tempdf.index, rotation=0)
486
+ else:
487
+ ax.xaxis.set_major_locator(plt.MaxNLocator(10))
488
+
489
+ # ax.xaxis.set_major_locator(plt.MaxNLocator(6))
490
+ ax.set_ylabel(name,fontsize='x-large')
491
+ # ax.set_xticklabels(tempdf.index.tolist(), rotation = 45,fontsize='x-large')
492
+ ## ax.xaxis.set_minor_locator(plt.NullLocator())
493
+ ## ax.tick_params(axis='x', labelleft=True)
494
+ # ax.xaxis.set_major_locator(years)
495
+ # ax.xaxis_date()
496
+ # ax.xaxis.set_major_formatter(years_fmt)
497
+ # ax.xaxis.set_minor_locator(months)
498
+ # ax.tick_params(axis='x', labelrotation=45,right = True)
499
+ if type(axvline) != type(None): axis.axvline(axvline)
500
+ fig.suptitle(ntitle,fontsize=20)
501
+ if 1:
502
+ fig.set_constrained_layout(True)
503
+
504
+ # plt.tight_layout()
505
+ # fig.subplots_adjust(top=top)
506
+ ...
507
+ return fig
508
+
509
+ #
510
+ def totexplain(self,pat='*',vtype='all',stacked=True,kind='bar',per='',top=0.9,title=''
511
+ ,use='level',threshold=0.0,ysize=10,**kwargs):
512
+ '''
513
+ Wrapper for different explanations
514
+ - :any:`explain_last`
515
+ - :any:`explain_per`
516
+ - :any:`explain_sum`
517
+ - :any:`explain_all`
518
+
519
+
520
+ Args:
521
+ pat (TYPE, optional): DESCRIPTION. Defaults to '*'.
522
+ vtype (per|all|last|sum, optional): what data to attribute. Defaults to 'all'.
523
+ stacked (TYPE, optional): DESCRIPTION. Defaults to True.
524
+ kind (TYPE, optional): DESCRIPTION. Defaults to 'bar'.
525
+ per (TYPE, optional): DESCRIPTION. Defaults to ''.
526
+ top (TYPE, optional): DESCRIPTION. Defaults to 0.9.
527
+ title (TYPE, optional): DESCRIPTION. Defaults to ''.
528
+ use (TYPE, optional): DESCRIPTION. Defaults to 'level'.
529
+ threshold (TYPE, optional): DESCRIPTION. Defaults to 0.0.
530
+ ysize (TYPE, optional): DESCRIPTION. Defaults to 10.
531
+ **kwargs (TYPE): DESCRIPTION.
532
+
533
+ Returns:
534
+ fig (TYPE): DESCRIPTION.
535
+
536
+ '''
537
+ if vtype.upper() == 'PER' :
538
+ fig = self.explain_per(pat=pat,per=per,top=top,use=use,title=title,threshold=threshold,ysize=ysize)
539
+
540
+ elif vtype.upper() == 'LAST' :
541
+ fig = self.explain_last(pat=pat,top=top,use=use,title=title,threshold=threshold)
542
+
543
+ elif vtype.upper() == 'SUM' :
544
+ fig = self.explain_sum(pat=pat,top=top,use=use,title=title,threshold=threshold)
545
+
546
+ else:
547
+ fig = self.explain_all(pat=pat,stacked=stacked,kind=kind,top=top,use=use,title=title,threshold=threshold)
548
+ return fig
549
+
550
+ # def get_att_gui(self,var='FY',spat = '*',desdic={},use='level'):
551
+ # '''Creates a jupyter ipywidget to display model level
552
+ # attributions '''
553
+ # def show_all2(Variable,Periode,Save,Use):
554
+ # global fig1,fig2
555
+ # fig1 = self.totexplain(pat=Variable,top=0.87,use=Use)
556
+ # fig2 = self.totexplain(pat=Variable,vtype='per',per = Periode,top=0.85,use=Use)
557
+ # if Save:
558
+ # fig1.savefig(f'Attribution-{Variable}-{use}.pdf')
559
+ # fig2.savefig(f'Attribution-{Variable}-{Periode}-{use}.pdf')
560
+ # print(f'Attribution-{Variable}-{use}.pdf and Attribution-{Variable}-{Periode}-{use}.pdf aare saved' )
561
+ #
562
+ # show = ip.interactive(show_all2,
563
+ # Variable = ip.Dropdown(options = sorted(self.model.endogene),value=var),
564
+ # Periode = self.model.current_per,
565
+ # Use = ip.RadioButtons(options= ['level', 'growth'],description='Use'),
566
+ # Save = False,
567
+ # )
568
+ # return show
569
+
570
+
571
+ if __name__ == '__main__' :
572
+ #%%
573
+ # running withe the mtotal model
574
+ df2 = pd.DataFrame({'Z':[1., 22., 33,43] , 'TY':[10.,20.,30.,40.] ,'YD':[10.,20.,30.,40.]},index=[2017,2018,2019,2020])
575
+ df3 = pd.DataFrame({'Z':[1., 22., 33,43] , 'TY':[10.,40.,60.,10.] ,'YD':[10.,49.,36.,40.]},index=[2017,2018,2019,2020])
576
+ ftest = '''
577
+ FRMl <> ii = TY(-1)+c(-1)+Z*c(-1) $
578
+ frml <> c=0.8*yd+log(1) $
579
+ frml <> d = c +2*ii(-1) $
580
+ frml <> c2=0.8*yd+log(1) $
581
+ frml <> d2 = c + 42*ii $
582
+ frml <> c3=0.8*yd+log(1) $
583
+ frml <> d3 = c +ii $
584
+ '''
585
+
586
+ m2=mc.model(ftest,straight=True,modelname='m2 testmodel')
587
+ df2=mc.insertModelVar(df2,m2)
588
+ df3=mc.insertModelVar(df3,m2)
589
+ z1 = m2(df2)
590
+ z2 = m2(df3)
591
+ ccc = m2.totexplain(pat='D2',per=2019,vtype='all',top=0.8)
592
+ ccc = m2.totexplain('D2',vtype='last',top=0.8)
593
+ ccc = m2.totexplain('D2',vtype='per',top=0.8)
594
+ #%%
595
+ ddd = totdif(m2)
596
+ eee = totdif(m2)
597
+ ddd.totexplain('D2',vtype='all',top=0.8,use='growth');
598
+ eee.totexplain('D2',vtype='all',top=0.8);
599
+
600
+ if False and ( not 'mtotal' in locals() ) :
601
+ # get the model
602
+ with open(r"models\mtotal.fru", "r") as text_file:
603
+ ftotal = text_file.read()
604
+
605
+ #get the data
606
+ base0 = pd.read_pickle(r'data\base0.pc')
607
+ base = pd.read_pickle(r'data\base.pc')
608
+ adve0 = pd.read_pickle(r'data\adve0.pc')
609
+ #%%
610
+ mtotal = mc.model(ftotal)
611
+
612
+ # prune(mtotal,base)
613
+ #%%
614
+ baseny = mtotal(base0 ,'2016q1','2018q4',samedata=False)
615
+ adverseny = mtotal(adve0 ,'2016q1','2018q4',samedata=True)
616
+ #%%
617
+ diff = mtotal.exodif() # exogeneous variables which are different between baseny and adverseny
618
+ #%%
619
+ assert 1==2 # just for stopping in test situations
620
+ #%%
621
+ adverseny = mtotal(adve0 ,'2016q1','2018q4',samedata=True) # to makew sure we have the right adverse.
622
+ countries = {c.split('__')[2] for c in diff.columns} # list of countries
623
+ countryexperiments = {e: [c for c in diff.columns if ('__'+e+'__') in c] for e in countries } # dic of experiments
624
+ assert len(diff.columns) == sum([len(c) for c in countryexperiments.values()]) , 'Not all exogeneous chocks variables are accountet for'
625
+ countryimpact = attribution(mtotal,countryexperiments,save='countryimpactxx',maxexp=30000,showtime = 1)
626
+ #%%
627
+ adverseny = mtotal(adve0 ,'2016q1','2018q4',samedata=True)
628
+ vartypes = {c.split('__') [1] for c in diff.columns}
629
+ vartypeexperiments = {e: [c for c in diff.columns if ('__'+e+'__') in c] for e in vartypes }
630
+ assert len(diff.columns) == sum([len(c) for c in vartypeexperiments.values()]) , 'Not all exogeneous chocks variables are accountet for'
631
+ vartypeimpact = attribution(mtotal,vartypeexperiments,save='vartypeimpactxx',maxexp=3000,showtime=1)
632
+ ##%%
633
+ #adverseny = mtotal(adve0 ,'2016q1','2018q4',samedata=True)
634
+ #allexo = {c[7:14] for c in diff.columns}
635
+ #allexoexperiments = {e: [c for c in diff.columns if ('__'+e+'__') in c] for e in allexo }
636
+ #allexoimpact = attribution(mtotal,allexoexperiments,base,adverseny,save='allexoimpact',maxexp=2000)
637
+ #%% test of upddf
638
+ if 0:
639
+ baseny = mtotal(base0 ,'2016q1','2018q4',samedata=False)
640
+ adverseny = mtotal(adve0 ,'2016q1','2018q4',samedata=True)
641
+ #%%
642
+ e = 'EE'
643
+ var = countryexperiments['EE']
644
+ vardiff = diff[var]
645
+ temp = adverseny[var].copy()
646
+ adverseny[var] = baseny[var]
647
+ temp2 = adverseny[var].copy()
648
+ _ = mc.upddf(adverseny,temp)
649
+ adverseny[var] = temp
650
+ temp3 = adverseny[var].copy()
651
+