style: 使用ruff 工具优化代码风格
This commit is contained in:
+21
-18
@@ -17,14 +17,14 @@ class RandomDataset(Dataset):
|
||||
self.length = length or int(np.random.randint(100, 200))
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {
|
||||
"input_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,))
|
||||
"target_ids": torch.randint(0, self.vocab_size, (self.max_length,)),
|
||||
}
|
||||
|
||||
|
||||
@@ -33,10 +33,10 @@ class MultiTurnDataset(Dataset):
|
||||
self.length = length or int(np.random.randint(100, 200))
|
||||
self.max_length = max_length
|
||||
self.vocab_size = vocab_size
|
||||
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
|
||||
def __getitem__(self, idx):
|
||||
input_ids = torch.randint(0, self.vocab_size, (self.max_length,))
|
||||
target_ids = torch.randint(0, self.vocab_size, (self.max_length,))
|
||||
@@ -54,18 +54,18 @@ class EarlyStoppingDataset(Dataset):
|
||||
self.length = length
|
||||
self.stop_after = stop_after
|
||||
self.count = 0
|
||||
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
|
||||
def __getitem__(self, idx):
|
||||
self.count += 1
|
||||
if self.count == self.stop_after:
|
||||
raise RuntimeError("Simulated early stopping")
|
||||
|
||||
|
||||
return {
|
||||
"input_ids": torch.randint(0, 1000, (64,)),
|
||||
"target_ids": torch.randint(0, 1000, (64,))
|
||||
"target_ids": torch.randint(0, 1000, (64,)),
|
||||
}
|
||||
|
||||
|
||||
@@ -74,10 +74,10 @@ def base_test_env(request: pytest.FixtureRequest):
|
||||
func_name = request.function.__name__
|
||||
test_dir = tempfile.mkdtemp(prefix=f"{func_name}_")
|
||||
config_path = os.path.join(test_dir, "config.json")
|
||||
|
||||
|
||||
n_dim_choices = [8, 16, 32]
|
||||
n_head_choices = [2, 4]
|
||||
|
||||
|
||||
dim = int(np.random.choice(n_dim_choices))
|
||||
n_heads = int(np.random.choice(n_head_choices))
|
||||
n_kv_heads = n_heads // 2
|
||||
@@ -91,16 +91,16 @@ def base_test_env(request: pytest.FixtureRequest):
|
||||
"dim_ffn": dim_ffn,
|
||||
"max_len": 1024,
|
||||
"n_layers": 4,
|
||||
"norm_eps": 1e-5
|
||||
"norm_eps": 1e-5,
|
||||
}
|
||||
|
||||
with open(config_path, 'w') as f:
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config, f)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
transformer_config = ModelConfig().load(config_path)
|
||||
model = Transformer(transformer_config).to(device=device)
|
||||
tokenizer = BpeTokenizer()
|
||||
|
||||
|
||||
yield {
|
||||
"device": device,
|
||||
"test_dir": str(test_dir),
|
||||
@@ -109,20 +109,23 @@ def base_test_env(request: pytest.FixtureRequest):
|
||||
"model": model,
|
||||
"tokenizer": tokenizer,
|
||||
}
|
||||
|
||||
|
||||
shutil.rmtree(test_dir)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def random_dataset():
|
||||
dataset = RandomDataset()
|
||||
yield dataset
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def multi_turn_dataset():
|
||||
dataset = MultiTurnDataset()
|
||||
yield dataset
|
||||
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def early_stopping_dataset():
|
||||
dataset = EarlyStoppingDataset()
|
||||
yield dataset
|
||||
yield dataset
|
||||
|
||||
Reference in New Issue
Block a user