Added more tests
This commit is contained in:
@@ -68,11 +68,11 @@ def split_data_by_month_percentage(df, percentages):
|
||||
tr, va, te = np.split(ids, [int((train_p/100) * len(ids)), int(((train_p + valid_p)/100) * len(ids))])
|
||||
return df.merge(tr, on=[year_str, month_str], how='inner'), df.merge(va, on=[year_str, month_str], how='inner'), df.merge(te, on=[year_str, month_str], how='inner')
|
||||
|
||||
def split_data_by_userdata_percentage(df, percentages):
|
||||
def split_data_by_userdata_percentage(df, percentages, sample):
|
||||
train_p, valid_p, test_p = percentages
|
||||
tr, va, te = pd.DataFrame(), pd.DataFrame(), pd.DataFrame()
|
||||
for user_id in df[user_str].unique():
|
||||
user_data = df[df[user_str]==user_id].sort_values([year_str, month_str])
|
||||
user_data = df[df[user_str]==user_id].sample(frac=sample/ 100).sort_values([year_str, month_str])
|
||||
u_tr, u_va, u_te = np.split(user_data, [int((train_p/100)*len(user_data)), int(((train_p+valid_p)/100)*len(user_data))])
|
||||
tr = pd.concat([tr, u_tr], ignore_index=True)
|
||||
va = pd.concat([va, u_va], ignore_index=True)
|
||||
@@ -285,9 +285,47 @@ def visualise_results_v2():
|
||||
# Fazit: keine eindeutig besseren Versionen erkennbar
|
||||
|
||||
|
||||
def test(model_type):
|
||||
sequence_length = 20
|
||||
data_filename = os.listdir(dataset_path)[0]
|
||||
timespan_id = hour_timespan_str
|
||||
threshold_id = with_threshold_str
|
||||
|
||||
file_path = os.path.join(dataset_path, data_filename)
|
||||
df = load_dataset(file_path)
|
||||
df = remove_covid_data(df)
|
||||
results = pd.DataFrame()
|
||||
|
||||
for percentage in [33,66,100]:
|
||||
print('Percentage:', percentage)
|
||||
tr,val,te = split_data_by_userdata_percentage(df, percentages=(80,10,10),sample=percentage)
|
||||
tr = reduce_columns(tr, data_filename)
|
||||
val = reduce_columns(val, data_filename)
|
||||
te = reduce_columns(te, data_filename)
|
||||
|
||||
user_data_train = prepare_user_data(tr)
|
||||
user_data_val = prepare_user_data(val)
|
||||
|
||||
best_model = train_models_v2(user_data_train, user_data_val,
|
||||
sequence_length=sequence_length,
|
||||
model_type=model_type)
|
||||
|
||||
results = pd.concat([results,
|
||||
evaluate_model_on_test_data(model=best_model,
|
||||
test_df=te,
|
||||
sequence_length=sequence_length,
|
||||
time_span_id=timespan_id,
|
||||
threshold_id=threshold_id,
|
||||
model_type=model_type,
|
||||
split_id=data_split_str)],
|
||||
ignore_index=True)
|
||||
print(results)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# main_two_v1()
|
||||
# visualise_results_v1()
|
||||
main_two_v2(model_type=model_type_gru)
|
||||
test(model_type=model_type_gru)
|
||||
# main_two_v2(model_type=model_type_gru)
|
||||
#visualise_results_v2()
|
||||
print('Done')
|
||||
|
||||
Reference in New Issue
Block a user