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.
modelhelp.py ADDED
@@ -0,0 +1,543 @@
1
+ """
2
+ Created on Tue Mar 7 10:38:28 2017
3
+
4
+ @author: hanseni
5
+
6
+ utilities for Stampe models
7
+
8
+ """
9
+
10
+
11
+ import networkx as nx
12
+ import pandas as pd
13
+ import numpy as np
14
+ import time
15
+ from contextlib import contextmanager
16
+ import sys
17
+ import itertools
18
+ import operator as op
19
+
20
+
21
+
22
+
23
+ def update_var(databank,xvar,operator='=',inputval=0,start='',end='',create=1, lprint=False,scale=1.0):
24
+
25
+ r"""Updates a variable in the databank. Possible update choices are:
26
+ \n \= : val = inputval
27
+ \n \+ : val = val + inputval
28
+ \n \- : val = val - inputval
29
+ \n \* : val = val * inputval
30
+ \n \=growth : val = val(t-1)+inputval +
31
+ \n \% : val = val(1+inputval/100)
32
+ \n
33
+ \n scale scales the input variables default =1.0
34
+
35
+ """
36
+ var = xvar.upper()
37
+ if var not in databank:
38
+ if not create:
39
+ errmsg = f'Variable to update not found:{var}, timespan = [{start} {end}] \nSet create=True if you want the variable created: '
40
+ raise Exception(errmsg)
41
+ else:
42
+ if 0:
43
+ print('Variable not in databank, created ',var)
44
+ databank[var]=0.0
45
+
46
+ current_per = databank.index[databank.index.get_slice_bound(start,'left')
47
+ :databank.index.get_slice_bound(end,'right')]
48
+ if operator.upper() == '+GROWTH':
49
+ if databank.index.get_loc(start) == 0:
50
+ raise Exception(f"+growth update can't start at first row, var:{var}")
51
+ orgdata=pd.Series(databank.pct_change().loc[current_per,var]).copy(deep=True)
52
+ ...
53
+ else:
54
+ orgdata=pd.Series(databank.loc[current_per,var]).copy(deep=True)
55
+
56
+ antalper=len(current_per)
57
+ # breakpoint()
58
+ if isinstance(inputval,float) or isinstance(inputval,int) :
59
+ inputliste=[float(inputval)]
60
+ elif isinstance(inputval,str):
61
+ inputliste=[float(i) for i in inputval.split()]
62
+ elif isinstance(inputval,list):
63
+ inputliste= [float(i) for i in inputval]
64
+
65
+
66
+ elif isinstance(inputval, pd.Series):
67
+ # inputliste= inputval.base
68
+ inputliste= list(inputval) #Ib for at håndtere mulitindex serier
69
+ else:
70
+ print('Fejl i inputdata',type(inputval))
71
+ inputdata=inputliste*antalper if len(inputliste) == 1 else inputliste
72
+
73
+ if len(inputdata) != antalper :
74
+ print('** Error, There should be',antalper,'values. There is:',len(inputdata))
75
+ print('** Update =',var,'Data=',inputdata,start,end)
76
+ raise Exception('wrong number of datapoints')
77
+ else:
78
+ inputserie=pd.Series(inputdata,current_per)*scale
79
+ # print(' Variabel------>',var)
80
+ # print( databank[var])
81
+ if operator=='=': #changes value to input value
82
+ outputserie=inputserie
83
+ elif operator == '+':
84
+ outputserie=orgdata+inputserie
85
+ elif operator == '*':
86
+ outputserie=orgdata*inputserie
87
+ elif operator == '%':
88
+ outputserie=orgdata*(1.0+inputserie/100.0)
89
+ elif operator == '=DIFF': # data=data(-1)+inputdata
90
+ if databank.index.get_loc(start) == 0:
91
+ raise Exception(f"=diff update can't start at first row, var:{var}")
92
+ ilocrow =databank.index.get_loc(start)-1
93
+ iloccol = databank.columns.get_loc(var)
94
+ temp=databank.iloc[ilocrow,iloccol]
95
+ addon = list(itertools.accumulate([i for i in inputserie],op.add))
96
+ # print(f'{addon=}')
97
+ opdater=[temp+add for add in addon]
98
+ outputserie=pd.Series(opdater,current_per)
99
+ elif operator.upper() == '=GROWTH': # data=data(-1)+inputdata
100
+ if databank.index.get_loc(start) == 0:
101
+ raise Exception(f"=growth update can't start at first row, var:{var}")
102
+ ilocrow =databank.index.get_loc(start)-1
103
+ iloccol = databank.columns.get_loc(var)
104
+ temp=databank.iloc[ilocrow,iloccol]
105
+ factor = list(itertools.accumulate([(1+i/100) for i in inputserie],op.mul))
106
+ opdater=[temp * it for it in factor]
107
+ outputserie=pd.Series(opdater,current_per)
108
+ elif operator.upper() == '+GROWTH': # data=data(-1)+inputdata
109
+ if databank.index.get_loc(start) == 0:
110
+ raise Exception(f"+growth update can't start at first row, var:{var}")
111
+ ilocrow =databank.index.get_loc(start)-1
112
+ iloccol = databank.columns.get_loc(var)
113
+ temp=databank.iloc[ilocrow,iloccol]
114
+ factor = list(itertools.accumulate([(1.0+i/100.0+o) for i,o in zip(inputserie,orgdata)],op.mul))
115
+ opdater=[temp * it for it in factor]
116
+ outputserie=pd.Series(opdater,current_per)
117
+ else:
118
+ raise Exception(f'Illegal operator in update:{operator} Variable: {var}')
119
+ outputserie=pd.Series(np.NaN,current_per)
120
+ outputserie.name=var
121
+ databank[var] = databank[var].astype('float') # to prevent error in the future
122
+ databank.loc[current_per,var]=outputserie.astype('float')
123
+
124
+ if lprint:
125
+ print('Update',operator,inputdata,start,end)
126
+ forspalte=str(max(6,len(var)))
127
+ print(('{:<'+forspalte+'} {:>20} {:>20} {:>20}').format(var,'Before', 'After', 'Diff'))
128
+ newdata=databank.loc[current_per,var]
129
+ diff=newdata-orgdata
130
+ for i in current_per:
131
+ print(('{:<'+forspalte+'} {:>20.4f} {:>20.4f} {:>20.4f}').format(str(i),orgdata[i],newdata[i],diff[i]))
132
+
133
+
134
+ def tovarlag(var,lag):
135
+ ''' creates a stringof var(lag) if lag else just lag '''
136
+ if type(lag)==int:
137
+ return f'{var}({lag:+})' if lag else var
138
+ else:
139
+ return f'{var}({lag})' if lag else var
140
+
141
+ def cutout(input,threshold=0.0):
142
+ '''get rid of rows below treshold and returns the dataframe or serie '''
143
+ if type(input)==pd.DataFrame:
144
+ org_sum = input.sum(axis=0)
145
+ new = input.iloc[(abs(input) >= threshold).any(axis=1).values,:]
146
+ if len(new) < len(input):
147
+ new_sum = new.sum(axis=0)
148
+ small = org_sum - new_sum
149
+ small.name = 'Small'
150
+ # breakpoint()
151
+ output = pd.concat([new,pd.DataFrame(small).T],axis=0)
152
+ else:
153
+ output = input
154
+ return output
155
+ if type(input)==pd.Series:
156
+ org_sum = input.sum()
157
+ new = input.iloc[(abs(input) >= threshold).values]
158
+ if len(new) < len(input):
159
+ new_sum = new.sum()
160
+ small = pd.Series(org_sum - new_sum)
161
+ small.index = ['Small']
162
+ output = pd.concat([new,small])
163
+ else:
164
+ output=input
165
+ return output
166
+
167
+ @contextmanager
168
+ def ttimer(input='test',show=True,short=False):
169
+ '''
170
+ A timer context manager, implemented using a
171
+ generator function. This one will report time even if an exception occurs"""
172
+
173
+ Parameters
174
+ ----------
175
+ input : string, optional
176
+ a name. The default is 'test'.
177
+ show : bool, optional
178
+ show the results. The default is True.
179
+ short : bool, optional
180
+ . The default is False.
181
+
182
+ Returns
183
+ -------
184
+ None.
185
+
186
+ '''
187
+
188
+ start = time.time()
189
+ if show and not short: print(f'{input} started at : {time.strftime("%H:%M:%S"):>{15}} ')
190
+ try:
191
+ yield
192
+ finally:
193
+ if show:
194
+ end = time.time()
195
+ seconds = (end - start)
196
+ minutes = seconds/60.
197
+ if minutes < 2.:
198
+ afterdec='1' if seconds >= 10 else ('3' if seconds >= 1 else '10')
199
+ print(f'{input} took : {seconds:>{15},.{afterdec}f} Seconds')
200
+ else:
201
+ afterdec='1' if minutes >= 10 else '4'
202
+ print(f'{input} took : {minutes:>{15},.{afterdec}f} Minutes')
203
+
204
+ def finddec(df):
205
+ ''' find a suitable number of decimal places from the magnitudes of a dataframe '''
206
+ try:
207
+ thismax = df.abs().max().max()
208
+ except:
209
+ thismax = df.abs().max()
210
+
211
+ if thismax > 1000:
212
+ outdec = 0
213
+ elif thismax > 10:
214
+ outdec = 3
215
+ else:
216
+ outdec = 6
217
+
218
+ return outdec
219
+
220
+
221
+ def insertModelVar(dataframe, model=None):
222
+ """Inserts all variables from model, not already in the dataframe.
223
+ Model can be a list of models """
224
+ if isinstance(model,list):
225
+ imodel=model
226
+ else:
227
+ imodel = [model]
228
+
229
+ myList=[]
230
+ for item in imodel:
231
+ myList.extend(item.allvar.keys())
232
+ manglervars = list(set(myList)-set(dataframe.columns))
233
+ if len(manglervars):
234
+ extradf = pd.DataFrame(0.0,index=dataframe.index,columns=manglervars).astype('float64')
235
+ data = pd.concat([dataframe,extradf],axis=1)
236
+ return data
237
+ else:
238
+ return dataframe
239
+
240
+ def df_extend(df,add=5):
241
+ '''Extends a Dataframe, assumes that the indes is of period_range type'''
242
+ newindex = pd.period_range(df.index[0], periods=len(df)+add, freq=df.index.freq)
243
+ return df.reindex(newindex,method='ffill')
244
+
245
+
246
+ import inspect
247
+ import ast
248
+ from pprint import pformat
249
+
250
+ def debug_var(*args, **kwargs):
251
+ """
252
+ Print names/expressions and values for args passed to debug_var,
253
+ with file, function, and line number. Robust to multi-line calls.
254
+ """
255
+ # --- Locate caller info ---
256
+ caller_frame = inspect.currentframe().f_back
257
+ func_name = caller_frame.f_code.co_name
258
+ file_name = caller_frame.f_code.co_filename
259
+ call_lineno = caller_frame.f_lineno
260
+
261
+ # --- Try to grab enough source to cover a potentially multi-line call ---
262
+ try:
263
+ # Entire function (or module) source + starting line
264
+ src_lines, start_line = inspect.getsourcelines(caller_frame)
265
+ rel_idx = call_lineno - start_line # index of the call line within src_lines
266
+
267
+ # Join forward from the call line until parentheses balance
268
+ # Find the first "debug_var(" occurrence on or after rel_idx
269
+ joined = "".join(src_lines[rel_idx:])
270
+ # If multiple statements share the line, trim before "debug_var("
271
+ start_pos = joined.find("debug_var(")
272
+ if start_pos == -1:
273
+ raise RuntimeError("Could not find 'debug_var(' on or after the call line.")
274
+
275
+ # Walk forward to capture the full call (balance parentheses)
276
+ i = start_pos
277
+ depth = 0
278
+ in_str = False
279
+ str_quote = ""
280
+ escaped = False
281
+ end_pos = None
282
+
283
+ while i < len(joined):
284
+ ch = joined[i]
285
+
286
+ if in_str:
287
+ if escaped:
288
+ escaped = False
289
+ elif ch == "\\":
290
+ escaped = True
291
+ elif ch == str_quote:
292
+ in_str = False
293
+ else:
294
+ if ch in ("'", '"'):
295
+ in_str = True
296
+ str_quote = ch
297
+ elif ch == "(":
298
+ depth += 1
299
+ elif ch == ")":
300
+ depth -= 1
301
+ if depth == 0:
302
+ end_pos = i + 1
303
+ break
304
+
305
+ i += 1
306
+
307
+ if end_pos is None:
308
+ raise RuntimeError("Could not balance parentheses for debug_var call.")
309
+
310
+ call_text = joined[start_pos:end_pos] # e.g., "debug_var(a, b=foo(x, y))"
311
+
312
+ # --- Parse the call with AST and extract argument source segments ---
313
+ # We parse the isolated call text so ast.get_source_segment works cleanly.
314
+ tree = ast.parse(call_text, mode="exec")
315
+ # Expecting one Expr(Call(...))
316
+ node = tree.body[0].value
317
+ if not isinstance(node, ast.Call):
318
+ raise RuntimeError("Parsed node was not a Call.")
319
+
320
+ # Helper to recover source of each arg from call_text
321
+ def src_of(subnode):
322
+ try:
323
+ return ast.get_source_segment(call_text, subnode)
324
+ except Exception:
325
+ return "<?>" # fallback
326
+
327
+ # Build argument name labels in the same order they were passed
328
+ # Positional args first:
329
+ pos_labels = [src_of(a) for a in node.args]
330
+ # Keyword args next:
331
+ kw_labels = [f"{src_of(k.arg) if k.arg else '<?>'}={src_of(k.value)}" for k in node.keywords if k.arg is not None]
332
+ # Handle **kwargs expansion specially (k.arg is None)
333
+ starstar_labels = [f"**{src_of(k.value)}" for k in node.keywords if k.arg is None]
334
+
335
+ # Combine labels to match runtime values ordering: positional, then explicit keywords
336
+ # Note: **kwargs expansion values appear only in source; we can't match them to individual items at runtime.
337
+ arg_labels = pos_labels + [lbl.split("=", 1)[0] for lbl in kw_labels] # use just the key before '='
338
+ # For display, we’ll list **expansions separately after the printed values.
339
+ starstar_note = ", ".join(starstar_labels) if starstar_labels else ""
340
+
341
+ except Exception:
342
+ # Fallback: source not available or parsing failed
343
+ arg_labels = ["<?>"] * len(args) + list(kwargs.keys())
344
+ starstar_note = ""
345
+
346
+ # --- Print header (location/context) ---
347
+ print(f"[debug_var] File: {file_name} | Function: {func_name} | Line: {call_lineno}")
348
+
349
+ # --- Print positional args ---
350
+ for label, value in zip(arg_labels[:len(args)], args):
351
+ if pd.DataFrame == type(value):
352
+ print(f"{label} = ")
353
+ print(f"{pformat(value)}")
354
+ else:
355
+ print(f" {label} = {pformat(value)}")
356
+
357
+ # --- Print keyword args (explicit ones only) ---
358
+ # Align labels for the kwargs we actually received (explicit, not **expansion members)
359
+ explicit_kw_keys = list(kwargs.keys())
360
+ kw_start = len(args)
361
+ for key in explicit_kw_keys:
362
+ # If parsing worked, try to use the parsed label; else fall back to the key
363
+ label = arg_labels[kw_start] if kw_start < len(arg_labels) else key
364
+ # Ensure the label matches the real key if we have a mismatch
365
+ if label != key:
366
+ label = key
367
+ print(f" {label} = {pformat(kwargs[key])}")
368
+ kw_start += 1
369
+
370
+ # --- Note any **kwargs expansions present in the source but not directly enumerable here ---
371
+ if starstar_note:
372
+ print(f" (source included {starstar_note})")
373
+
374
+ def colab_link(notebook, folder='simulation',
375
+ user='IbHansen', repo='wb-debt-simulation', branch='main', badge=False, render=True):
376
+ """Display or print a Google Colab link for a Jupyter notebook hosted on GitHub.
377
+
378
+ Args:
379
+ notebook: Notebook name without .ipynb extension, e.g. 'my_notebook'.
380
+ folder: Repository folder containing the notebook.
381
+ user: GitHub username.
382
+ repo: GitHub repository name.
383
+ branch: Git branch.
384
+ badge: If True, use the 'Open in Colab' badge image instead of a plain URL.
385
+ render: If True and badge=True, render the badge in the notebook.
386
+ If False, print the raw HTML string for copy-pasting.
387
+
388
+ Returns:
389
+ None. Output is displayed or printed.
390
+ """
391
+ url = f'https://colab.research.google.com/github/{user}/{repo}/blob/{branch}/{folder}/{notebook}.ipynb'
392
+ if badge:
393
+ badge_html = f'<a href="{url}" target="_blank"><img src="https://colab.research.google.com/assets/colab-badge.svg" alt="Open In Colab"/></a>'
394
+ if render:
395
+ from IPython.display import HTML, display
396
+ display(HTML(badge_html))
397
+ else:
398
+ print(badge_html)
399
+ else:
400
+ print(url)
401
+
402
+ def build_sorted_bond_desc_dict(var_names):
403
+ """
404
+ Build an ordered dictionary mapping variable names to human-readable bond descriptions.
405
+
406
+ The function expects variable names to follow a naming convention where the
407
+ relevant bond information appears after a double underscore, for example:
408
+
409
+ PREFIX__10_YEAR_DOM
410
+ PREFIX__5_YEAR_USD
411
+
412
+ From this suffix, the function extracts:
413
+ - maturity: the first element, interpreted as an integer number of years
414
+ - currency: the third element, expected to be ``DOM`` or another currency code
415
+
416
+ It then creates descriptions such as:
417
+ - ``"10 year domestic bond"``
418
+ - ``"5 year USD bond"``
419
+
420
+ The returned dictionary is sorted first by currency, with domestic bonds
421
+ appearing before foreign-currency bonds, and then by maturity in ascending order.
422
+
423
+ Parameters
424
+ ----------
425
+ var_names : iterable of str
426
+ Collection of variable names formatted like ``PREFIX__<maturity>_YEAR_<currency>``.
427
+
428
+ Returns
429
+ -------
430
+ dict
431
+ An ordered dictionary-like mapping from each variable name to its
432
+ generated description. In Python 3.7 and later, the standard ``dict``
433
+ preserves insertion order.
434
+
435
+ Raises
436
+ ------
437
+ IndexError
438
+ If a variable name does not contain the expected ``"__"`` separator or
439
+ does not have enough underscore-separated parts after it.
440
+ ValueError
441
+ If the maturity part cannot be converted to an integer.
442
+
443
+ Example
444
+ -------
445
+ >>> build_sorted_bond_desc_dict(["BOND__10_YEAR_DOM", "BOND__5_YEAR_USD", "BOND__2_YEAR_DOM"])
446
+ {
447
+ "BOND__2_YEAR_DOM": "2 year domestic bond",
448
+ "BOND__10_YEAR_DOM": "10 year domestic bond",
449
+ "BOND__5_YEAR_USD": "5 year USD bond"
450
+ }
451
+ """
452
+ parsed = []
453
+
454
+ for v in var_names:
455
+ tail = v.split('__')[1] # e.g. '10_YEAR_DOM'
456
+ parts = tail.split('_') # ['10', 'YEAR', 'DOM']
457
+
458
+ maturity = int(parts[0]) # numeric for sorting
459
+ currency = parts[2].lower() # dom or usd
460
+
461
+ currency_label = 'domestic' if currency == 'dom' else currency.upper()
462
+ description = f"{maturity} year {currency_label} bond"
463
+
464
+ # sort key: domestic first (0), then usd (1), then maturity
465
+ currency_order = 0 if currency == 'dom' else 1
466
+
467
+ parsed.append((currency_order, maturity, v, description))
468
+
469
+ # Sort by currency first, then maturity
470
+ parsed.sort()
471
+
472
+ # Build ordered dict (normal dict preserves order in Python 3.7+)
473
+ return {v: desc for _, _, v, desc in parsed}
474
+
475
+ def build_sorted_rate_desc_dict(var_names):
476
+ """
477
+ Build an ordered dictionary mapping interest-rate variable names to
478
+ human-readable bond interest rate descriptions.
479
+
480
+ Variable names are expected to end with ``__<maturity>``. If the part
481
+ before the final separator ends with ``_<CCY>``, where ``<CCY>`` is a
482
+ 3-letter alphabetic currency mnemonic, that mnemonic is used as the
483
+ currency label. Otherwise, the bond is treated as domestic-currency
484
+ denominated.
485
+
486
+ The result is sorted with domestic bonds first, then foreign-currency
487
+ bonds alphabetically by currency mnemonic, and then by maturity.
488
+
489
+ Parameters
490
+ ----------
491
+ var_names : iterable of str
492
+ Variable names such as ``interest_rate__10``,
493
+ ``interest_rate_USD__5``, or ``interest_rate_EUR__2``.
494
+
495
+ Returns
496
+ -------
497
+ dict
498
+ Ordered mapping from variable names to descriptions.
499
+
500
+ Raises
501
+ ------
502
+ IndexError
503
+ If a variable name does not contain ``"__"``.
504
+ ValueError
505
+ If the maturity suffix cannot be parsed as an integer.
506
+ """
507
+ parsed = []
508
+
509
+ for v in var_names:
510
+ left, maturity = v.rsplit('__', 1)
511
+ maturity = int(maturity)
512
+
513
+ suffix = left.rsplit('_', 1)[-1]
514
+ if len(suffix) == 3 and suffix.isalpha():
515
+ currency = suffix.upper()
516
+ currency_order = 1
517
+ currency_label = currency
518
+ else:
519
+ currency = ''
520
+ currency_order = 0
521
+ currency_label = 'domestic'
522
+
523
+ description = f"{maturity} year {currency_label} bond interest rate"
524
+ parsed.append((currency_order, currency, maturity, v, description))
525
+
526
+ parsed.sort()
527
+
528
+ return {v: desc for _, _, _, v, desc in parsed}
529
+
530
+
531
+ if __name__ == '__main__':
532
+ #%% Test
533
+ if not 'baseline' in locals() or 1 :
534
+ from modelclass import model
535
+ madam,baseline = model.modelload('../Examples/ADAM/baseline.pcim',run=1,silent=0,ljit=1,stringjit=0 )
536
+ # make a simpel experimet VAT
537
+ scenarie = baseline.copy()
538
+ scenarie.TG = scenarie.TG + 0.05
539
+ _ = madam(scenarie)
540
+
541
+ update_var(baseline,'tg','=GROWTH',1,2021,2024,lprint=1)
542
+
543
+