(path: str)
| 119 | |
| 120 | |
| 121 | def load_cases(path: str) -> List[Phase7Case]: |
| 122 | cases: List[Phase7Case] = [] |
| 123 | with open(path, "r", encoding="utf-8") as f: |
| 124 | for line in f: |
| 125 | line = line.rstrip("\n") |
| 126 | if not line: |
| 127 | continue |
| 128 | fields = line.split("\t") |
| 129 | while len(fields) < 5: |
| 130 | fields.append("") |
| 131 | mem_cases = fields[2].split("|") if fields[2] else [] |
| 132 | spatial_tpos = parse_int_list(fields[3]) |
| 133 | ptr_tpos = parse_int_list(fields[4]) |
| 134 | assert len(mem_cases) == len(spatial_tpos) == len(ptr_tpos) |
| 135 | cases.append( |
| 136 | Phase7Case( |
| 137 | case_id=fields[0], |
| 138 | label=fields[1], |
| 139 | mem_cases=mem_cases, |
| 140 | spatial_tpos=spatial_tpos, |
| 141 | ptr_tpos=ptr_tpos, |
| 142 | ) |
| 143 | ) |
| 144 | return cases |
| 145 | |
| 146 | |
| 147 | def layer_norm_2d(x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor, eps: float = 1e-6) -> torch.Tensor: |
no test coverage detected