Team Ai
Apppublic

h3nock/scriptify-api

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
inference_utils.py89 linesDownload Raw Back to root
1from pathlib import Path2from typing import Dict, NamedTuple, Union3import numpy as np4import torch 5 6NULL_CHAR = '\x00'7 8 9class PrimingData(NamedTuple):10    """combines data required for priming the HandwritingRNN sampling"""11    stroke_tensors: torch.Tensor # (batch_size, num_prime_strokes, 3) 12    char_seq_tensors: torch.Tensor # (batch_size, num_prime_chars)13    char_seq_lengths: torch.Tensor # (batch_size,)14 15def construct_alphabet_list(alphabet_string: str) -> list[str]:16    if not isinstance(alphabet_string, str):17        raise TypeError("alphabet_string must be a string") 18    19    char_list = list(alphabet_string) 20    return [NULL_CHAR] + char_list 21 22def get_alphabet_map(alphabet_list: list[str]) -> Dict[str, int]:23    """creates a char to index map from full alphabet list"""24    return {char: idx for idx, char in enumerate(alphabet_list)}  25 26def encode_text(text: str, char_to_index_map: Dict[str, int], 27                max_length: int, add_eos: bool = True, eos_char_index: int = 028                ) -> tuple[np.ndarray, int]:29    """Encode a text string into a sequence of integer indices"""30    encoded = [char_to_index_map.get(c, eos_char_index) for c in text] 31    if add_eos:32        encoded.append(eos_char_index) 33 34    true_length = len(encoded)35 36    if true_length <= max_length: 37        padded_encoded = np.full(max_length, eos_char_index, dtype=np.int64) 38        padded_encoded[:true_length] = encoded 39    else:40        padded_encoded = np.array(encoded[:max_length], dtype=np.int64) 41        true_length = max_length 42    43    return np.array([padded_encoded]), true_length44 45 46def convert_offsets_to_absolute_coords(stroke_offsets: list[list[float]]) -> list[list[float]]:47    if not stroke_offsets:48        return []49    50    # convert to numpy for vectorized operations51    strokes_array = np.array(stroke_offsets)52    53    # vectorized cumulative sum for x and y 54    strokes_array[:, 0] = np.cumsum(strokes_array[:, 0])  # cumulative dx55    strokes_array[:, 1] = np.cumsum(strokes_array[:, 1])  # cumulative dy56    57    return strokes_array.tolist()58 59 60def load_np_strokes(stroke_path: Union[Path, str]) -> np.ndarray:61    """loads stroke sequence from stroke_path"""62    stroke_path = Path(stroke_path)63    if not stroke_path.exists():64        raise FileNotFoundError(f"style strokes file not found at {stroke_path}")65    66    return np.load(stroke_path)67 68def load_text(text_path: Union[Path, str]) -> str:69    """loads text from a text_path""" 70    text_path = Path(text_path) 71    if not text_path.exists():72        raise FileNotFoundError(f"Text file not found at {text_path}")73    if not text_path.is_file():74        raise IsADirectoryError(f"Path is a directory, not a file.")75    76    try: 77        with open(text_path, 'r', encoding='utf-8') as f:78            content = f.read() 79        return content 80 81    except Exception as e:82        raise IOError(f"Error reading text file {text_path}: {e}")83 84def load_priming_data(style: int):85    86    priming_text = load_text(f"./styles/style{style}.txt")87    priming_strokes = load_np_strokes(f"./styles/style{style}.npy")88    89    return priming_text, priming_strokes