You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

utils.py 4.6 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109
  1. import os
  2. from collections import defaultdict
  3. import json
  4. import matplotlib.pyplot as plt
  5. import numpy as np
  6. from loguru import logger
  7. import traceback
  8. from examples.dataset_table_workflow.config import *
  9. from benchmarks.config import default_size_list
  10. class Recorder:
  11. def __init__(self, headers=["Mean", "Std Dev"], formats=["{:.2f}", "{:.2f}"]):
  12. assert len(headers) == len(formats), "Headers and formats length must match."
  13. self.data = defaultdict(lambda: defaultdict(list))
  14. self.headers = headers
  15. self.formats = formats
  16. def record(self, user, idx, scores):
  17. self.data[user][idx].append(scores)
  18. def get_performance_data(self, user):
  19. if user in self.data:
  20. return [idx_scores for idx_scores in self.data[user].values()]
  21. else:
  22. return []
  23. def save(self, path):
  24. with open(path, "w") as f:
  25. json.dump(self.data, f, indent=4, default=list)
  26. def load(self, path):
  27. with open(path, "r") as f:
  28. self.data = json.load(f, object_hook=lambda x: defaultdict(list, x))
  29. def should_test_method(self, user, idx, path):
  30. if os.path.exists(path):
  31. self.load(path)
  32. return user not in self.data or str(idx) not in self.data[user]
  33. return True
  34. def process_single_aug(user, idx, scores, recorders, root_path):
  35. try:
  36. scores_array = np.array(scores)
  37. while scores_array.ndim < 3:
  38. scores_array = scores_array[np.newaxis, :]
  39. select_scores = scores_array[:, 0, :].tolist()
  40. mean_scores = np.mean(scores_array, axis=1).tolist()
  41. oracle_scores = np.min(scores_array, axis=1).tolist()
  42. for method_name, scores in zip(["select_score", "mean_score", "oracle_score"],
  43. [select_scores, mean_scores, oracle_scores]):
  44. recorders[method_name].record(user, idx, scores)
  45. save_path = os.path.join(root_path, f"{method_name}_performance.json")
  46. recorders[method_name].save(save_path)
  47. except Exception as e:
  48. error_message = traceback.format_exc()
  49. logger.error(f"Error in process_single_aug for user {user}, idx {idx}: {error_message}")
  50. def analyze_performance(user, recorders):
  51. oracle_score_list = recorders["hetero_oracle_score"].get_performance_data(user)
  52. select_score_list = recorders["hetero_select_score"].get_performance_data(user)
  53. multi_avg_score_list = recorders["hetero_multiple_avg"].get_performance_data(user)
  54. mean_differences = {}
  55. for user_id in range(len(oracle_score_list)):
  56. select_scores = select_score_list[user_id]
  57. oracle_scores = oracle_score_list[user_id]
  58. mean_difference = np.mean(select_scores) - np.mean(oracle_scores)
  59. mean_differences[user_id] = mean_difference
  60. sorted_user_ids = sorted(mean_differences, key=mean_differences.get, reverse=True)
  61. for user_id in sorted_user_ids:
  62. single_multi_diff = np.mean(select_score_list[user_id]) - np.mean(multi_avg_score_list[user_id])
  63. logger.info(f"{user}, {user_id}, {mean_differences[user_id]}, {single_multi_diff}")
  64. def plot_performance_curves(user, recorders, task="Hetero", n_labeled_list=default_size_list):
  65. plt.figure(figsize=(10, 6))
  66. for method, recorder in recorders.items():
  67. if method == "hetero_single_aug":
  68. continue
  69. user_data = recorder.get_performance_data(user)
  70. if user_data:
  71. scores_array = np.array([np.array(lst) for lst in user_data])
  72. mean_scores = np.squeeze(np.mean(scores_array, axis=0))
  73. std_scores = np.squeeze(np.std(scores_array, axis=0))
  74. method_plot = '_'.join(method.split('_')[1:]) if method not in ['user_model', 'oracle_score', 'select_score', 'mean_score'] else method
  75. style = styles.get(method_plot, {"color": "black", "linestyle": "-"})
  76. plt.plot(range(len(n_labeled_list)), mean_scores, label=labels.get(method_plot), **style)
  77. std_scale = 0.2 if task == "Hetero" else 0.5
  78. plt.fill_between(range(len(n_labeled_list)), mean_scores - std_scale * std_scores, mean_scores + std_scale * std_scores, color=style["color"], alpha=0.2)
  79. plt.xticks(range(len(n_labeled_list)), n_labeled_list)
  80. plt.xlabel('Sample Size')
  81. plt.ylabel('RMSE')
  82. plt.title(f'Table {task} Limited Labeled Data')
  83. plt.legend()
  84. plt.tight_layout()
  85. plt.savefig(os.path.join('./results/figs', f"{user}_labeled_{list(recorders.keys())}.png"), bbox_inches="tight", dpi=700)