refine the graph model interface (#136)

This commit is contained in:
Xu Yang
2024-07-30 19:38:37 +08:00
committed by GitHub
parent f6dc1f4d3f
commit 82a2a5d5b4
4 changed files with 26 additions and 11 deletions
+17 -8
View File
@@ -1,6 +1,7 @@
import json
import pickle
import site
import traceback
import uuid
from pathlib import Path
from typing import Any, Dict, Optional
@@ -78,8 +79,7 @@ class ModelFBWorkspace(FBWorkspace):
batch_size: int = 8,
num_features: int = 10,
num_timesteps: int = 4,
num_nodes: int = 50,
num_edges: int = 100,
num_edges: int = 20,
input_value: float = 1.0,
param_init_value: float = 1.0,
):
@@ -106,9 +106,9 @@ class ModelFBWorkspace(FBWorkspace):
m = model_cls(num_features=input_shape[1], num_timesteps=input_shape[2])
data = torch.full(input_shape, input_value)
elif self.target_task.model_type == "Graph":
node_feature = torch.randn(batch_size, num_nodes, num_features)
edge_index = torch.randint(0, num_nodes, (2, num_edges))
m = model_cls(num_nodes=num_nodes, num_features=num_features)
node_feature = torch.randn(batch_size, num_features)
edge_index = torch.randint(0, batch_size, (2, num_edges))
m = model_cls(num_features=num_features)
data = (node_feature, edge_index)
else:
raise ValueError(f"Unsupported model type: {self.target_task.model_type}")
@@ -129,10 +129,19 @@ class ModelFBWorkspace(FBWorkspace):
if MODEL_IMPL_SETTINGS.enable_execution_cache:
pickle.dump((execution_feedback_str, execution_model_output), open(cache_file_path, "wb"))
return execution_feedback_str, execution_model_output
except Exception as e:
return f"Execution error: {e}", None
execution_feedback_str = f"Execution error: {e}\nTraceback: {traceback.format_exc()}"
execution_model_output = None
code_path = self.workspace_path / f"model.py"
execution_feedback_str = execution_feedback_str.replace(str(code_path.parent.absolute()), r"/path/to").replace(
str(site.getsitepackages()[0]), r"/path/to/site-packages"
)
if len(execution_feedback_str) > 2000:
execution_feedback_str = (
execution_feedback_str[:1000] + "....hidden long error message...." + execution_feedback_str[-1000:]
)
return execution_feedback_str, execution_model_output
ModelExperiment = Experiment
@@ -106,6 +106,8 @@ evaluator_code_feedback:
Notice that your critics are not for user to debug the code. They are sent to the coding agent to correct the code. So don't give any following items for the user to check like "Please check the code line XXX".
You suggestion should not include any code, just some clear and short suggestions. Please point out very critical issues in your response, ignore non-important issues to avoid confusion. If no big issue found in the code, you can response "No critics found".
You should provide the suggestion to each of your critic to help the user improve the code. Please response the critic in the following format. Here is an example structure for the output:
critic 1: The critic message to critic 1
critic 2: The critic message to critic 2