Added new evaluation method
This commit is contained in:
+14
-19
@@ -54,6 +54,18 @@ def prepare_user_data(df):
|
||||
users = df_sorted['user'].unique()
|
||||
return {user: df_sorted[df_sorted['user'] == user] for user in users}
|
||||
|
||||
def prepare_data_for_model(user_data, sequence_length):
|
||||
X, y = [], []
|
||||
for user, data in user_data.items():
|
||||
features = data.drop('user', axis=1).values
|
||||
labels = data['user'].values
|
||||
for i in range(len(features) - sequence_length):
|
||||
X.append(features[i:i + sequence_length])
|
||||
y.append(labels[i + sequence_length])
|
||||
X = np.array(X)
|
||||
y = np.array(y)
|
||||
return X,y
|
||||
|
||||
# === Training & Validation ===
|
||||
def train_models(user_data, user_data_val, sequence_lengths=[20], tuner_dir="./working/tuner"):
|
||||
best_models = {}
|
||||
@@ -65,25 +77,8 @@ def train_models(user_data, user_data_val, sequence_lengths=[20], tuner_dir="./w
|
||||
|
||||
for sequence_length in sequence_lengths:
|
||||
print(f"\n=== Training for Sequence Length: {sequence_length} ===")
|
||||
X, y = [], []
|
||||
for user, data in user_data.items():
|
||||
features = data.drop('user', axis=1).values
|
||||
labels = data['user'].values
|
||||
for i in range(len(features) - sequence_length):
|
||||
X.append(features[i:i + sequence_length])
|
||||
y.append(labels[i + sequence_length])
|
||||
X = np.array(X)
|
||||
y = np.array(y)
|
||||
|
||||
X_val, y_val = [], []
|
||||
for user, data in user_data_val.items():
|
||||
features = data.drop('user', axis=1).values
|
||||
labels = data['user'].values
|
||||
for i in range(len(features) - sequence_length):
|
||||
X_val.append(features[i:i + sequence_length])
|
||||
y_val.append(labels[i + sequence_length])
|
||||
X_val = np.array(X_val)
|
||||
y_val = np.array(y_val)
|
||||
X, y = prepare_data_for_model(user_data=user_data, sequence_length=sequence_length)
|
||||
X_val, y_val = prepare_data_for_model(user_data=user_data_val, sequence_length=sequence_length)
|
||||
|
||||
if X.shape[0] == 0 or X_val.shape[0] == 0:
|
||||
print(f"⚠️ Skipped sequence length {sequence_length} due to insufficient data.")
|
||||
|
||||
Reference in New Issue
Block a user